@@ -1138,7 +1138,7 @@ static FORCE_INLINE float L2BinaryFloatDistanceFromHammingDist ( const Binary4Bi
11381138
11391139
11401140template <bool BUILD , int64_t (*DOTPRODUCT_FN )( const uint8_t * pVec4Bit, const uint8_t * pVec1Bit, int iBytes )>
1141- static float IPBinaryFloatDistance ( const void * __restrict pVect1, const void * __restrict pVect2, size_t uRowID1, size_t uRowID2, const void * __restrict pParam )
1141+ FORCE_INLINE static float IPBinaryFloatDistance ( const void * __restrict pVect1, const void * __restrict pVect2, size_t uRowID1, size_t uRowID2, const void * __restrict pParam )
11421142{
11431143 const auto & tBinaryParam = *(const DistFuncParamBinary_t*)pParam;
11441144
@@ -1173,7 +1173,7 @@ static float IPBinaryFloatDistance ( const void * __restrict pVect1, const void
11731173// in org.elasticsearch.index.codec.vectors.es816.ES816BinaryFlatVectorsScorer
11741174// Permalink: https://github.com/elastic/elasticsearch/blob/1dd41ec2b683a7b7c9c16af404b842cf85cbd5bc/server/src/main/java/org/elasticsearch/index/codec/vectors/es816/ES816BinaryFlatVectorsScorer.java
11751175template <bool BUILD , int64_t (*DOTPRODUCT_FN )( const uint8_t * pVec4Bit, const uint8_t * pVec1Bit, int iBytes )>
1176- static float L2BinaryFloatDistance ( const void * __restrict pVect1, const void * __restrict pVect2, size_t uRowID1, size_t uRowID2, const void * __restrict pParam )
1176+ FORCE_INLINE static float L2BinaryFloatDistance ( const void * __restrict pVect1, const void * __restrict pVect2, size_t uRowID1, size_t uRowID2, const void * __restrict pParam )
11771177{
11781178 const auto & tBinaryParam = *(const DistFuncParamBinary_t*)pParam;
11791179
@@ -1211,32 +1211,60 @@ float IPBinaryFloatDistanceGeneric ( const void * pVect1, const void * pVect2, s
12111211 return IPBinaryFloatDistance<false ,BinaryDotProduct> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
12121212}
12131213
1214+ float IPBinaryFloatDistanceGenericBuild ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
1215+ {
1216+ return IPBinaryFloatDistance<true ,BinaryDotProduct> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
1217+ }
1218+
12141219float L2BinaryFloatDistanceGeneric ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
12151220{
12161221 return L2BinaryFloatDistance<false ,BinaryDotProduct> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
12171222}
12181223
1224+ float L2BinaryFloatDistanceGenericBuild ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
1225+ {
1226+ return L2BinaryFloatDistance<true ,BinaryDotProduct> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
1227+ }
1228+
12191229#if !defined(USE_SIMDE)
12201230
12211231float IPBinaryFloatDistanceSIMD16 ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
12221232{
12231233 return IPBinaryFloatDistance<false ,BinaryDotProduct16<false >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
12241234}
12251235
1236+ float IPBinaryFloatDistanceSIMD16Build ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
1237+ {
1238+ return IPBinaryFloatDistance<true ,BinaryDotProduct16<false >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
1239+ }
1240+
12261241float IPBinaryFloatDistanceSIMD16Residuals ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
12271242{
12281243 return IPBinaryFloatDistance<false ,BinaryDotProduct16<true >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
12291244}
12301245
1231- template <bool RESIDUALS >
1232- static void IPBinaryFloatDistanceBatch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t , size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
1246+ float IPBinaryFloatDistanceSIMD16ResidualsBuild ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
1247+ {
1248+ return IPBinaryFloatDistance<true ,BinaryDotProduct16<true >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
1249+ }
1250+
1251+ template <bool BUILD , bool RESIDUALS >
1252+ FORCE_INLINE static void IPBinaryFloatDistanceBatch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
12331253{
12341254 const auto & tBinaryParam = *(const DistFuncParamBinary_t*)pParam;
12351255
12361256 auto pV1 = (const uint8_t *)pVect1;
12371257 auto pVA = (const uint8_t *)pVect2A;
12381258 auto pVB = (const uint8_t *)pVect2B;
12391259
1260+ // Build mode: source is 1-bit raw data; fetch its 4-bit representation from the pool
1261+ // using the source row id. Amortized over both candidates by doing this once.
1262+ if constexpr ( BUILD )
1263+ {
1264+ if ( uRowID1!=(size_t )-1 )
1265+ pV1 = tBinaryParam.m_fnFetcher (uRowID1);
1266+ }
1267+
12401268 assert ( uRowID2A!=(size_t )-1 );
12411269 assert ( uRowID2B!=(size_t )-1 );
12421270
@@ -1258,33 +1286,59 @@ static void IPBinaryFloatDistanceBatch2 ( const void * pVect1, const void * pVec
12581286
12591287void IPBinaryFloatDistanceSIMD16Batch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
12601288{
1261- IPBinaryFloatDistanceBatch2<false > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1289+ IPBinaryFloatDistanceBatch2<false ,false > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1290+ }
1291+
1292+ void IPBinaryFloatDistanceSIMD16Batch2Build ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
1293+ {
1294+ IPBinaryFloatDistanceBatch2<true ,false > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
12621295}
12631296
12641297void IPBinaryFloatDistanceSIMD16ResidualsBatch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
12651298{
1266- IPBinaryFloatDistanceBatch2<true > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1299+ IPBinaryFloatDistanceBatch2<false ,true > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1300+ }
1301+
1302+ void IPBinaryFloatDistanceSIMD16ResidualsBatch2Build ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
1303+ {
1304+ IPBinaryFloatDistanceBatch2<true ,true > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
12671305}
12681306
12691307float L2BinaryFloatDistanceSIMD16 ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
12701308{
12711309 return L2BinaryFloatDistance<false ,BinaryDotProduct16<false >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
12721310}
12731311
1312+ float L2BinaryFloatDistanceSIMD16Build ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
1313+ {
1314+ return L2BinaryFloatDistance<true ,BinaryDotProduct16<false >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
1315+ }
1316+
12741317float L2BinaryFloatDistanceSIMD16Residuals ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
12751318{
12761319 return L2BinaryFloatDistance<false ,BinaryDotProduct16<true >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
12771320}
12781321
1279- template <bool RESIDUALS >
1280- static void L2BinaryFloatDistanceBatch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t , size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
1322+ float L2BinaryFloatDistanceSIMD16ResidualsBuild ( const void * pVect1, const void * pVect2, size_t uRowID1, size_t uRowID2, const void * pParam )
1323+ {
1324+ return L2BinaryFloatDistance<true ,BinaryDotProduct16<true >> ( pVect1, pVect2, uRowID1, uRowID2, pParam );
1325+ }
1326+
1327+ template <bool BUILD , bool RESIDUALS >
1328+ static FORCE_INLINE void L2BinaryFloatDistanceBatch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
12811329{
12821330 const auto & tBinaryParam = *(const DistFuncParamBinary_t*)pParam;
12831331
12841332 auto pV1 = (const uint8_t *)pVect1;
12851333 auto pVA = (const uint8_t *)pVect2A;
12861334 auto pVB = (const uint8_t *)pVect2B;
12871335
1336+ if constexpr ( BUILD )
1337+ {
1338+ if ( uRowID1!=(size_t )-1 )
1339+ pV1 = tBinaryParam.m_fnFetcher (uRowID1);
1340+ }
1341+
12881342 assert ( uRowID2A!=(size_t )-1 );
12891343 assert ( uRowID2B!=(size_t )-1 );
12901344
@@ -1306,12 +1360,22 @@ static void L2BinaryFloatDistanceBatch2 ( const void * pVect1, const void * pVec
13061360
13071361void L2BinaryFloatDistanceSIMD16Batch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
13081362{
1309- L2BinaryFloatDistanceBatch2<false > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1363+ L2BinaryFloatDistanceBatch2<false ,false > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1364+ }
1365+
1366+ void L2BinaryFloatDistanceSIMD16Batch2Build ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
1367+ {
1368+ L2BinaryFloatDistanceBatch2<true ,false > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
13101369}
13111370
13121371void L2BinaryFloatDistanceSIMD16ResidualsBatch2 ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
13131372{
1314- L2BinaryFloatDistanceBatch2<true > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1373+ L2BinaryFloatDistanceBatch2<false ,true > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
1374+ }
1375+
1376+ void L2BinaryFloatDistanceSIMD16ResidualsBatch2Build ( const void * pVect1, const void * pVect2A, const void * pVect2B, size_t uRowID1, size_t uRowID2A, size_t uRowID2B, const void * pParam, float & fDistA , float & fDistB )
1377+ {
1378+ L2BinaryFloatDistanceBatch2<true ,true > ( pVect1, pVect2A, pVect2B, uRowID1, uRowID2A, uRowID2B, pParam, fDistA , fDistB );
13151379}
13161380
13171381#endif // !USE_SIMDE
@@ -1341,13 +1405,10 @@ IPSpaceBinaryFloat_c::IPSpaceBinaryFloat_c ( size_t uDim, bool bBuild )
13411405 , m_tDistFuncParam ( uDim )
13421406{
13431407#if defined(USE_SIMDE)
1344- if ( bBuild )
1345- m_fnDist = IPBinaryFloatDistance<true ,BinaryDotProduct>;
1346- else
1347- {
1348- m_fnDist = IPBinaryFloatDistance<false ,BinaryDotProduct>;
1349- m_eDistFuncId = DistFuncId_e::IP_BINARY_GENERIC ;
1350- }
1408+ m_fnDist = bBuild
1409+ ? IPBinaryFloatDistance<true ,BinaryDotProduct>
1410+ : IPBinaryFloatDistance<false ,BinaryDotProduct>;
1411+ m_eDistFuncId = DistFuncId_e::IP_BINARY_GENERIC ;
13511412#else
13521413 int iBytes = ( uDim+7 ) >> 3 ;
13531414 bool bUseSSE = iBytes>=16 ;
@@ -1367,13 +1428,10 @@ IPSpaceBinaryFloat_c::IPSpaceBinaryFloat_c ( size_t uDim, bool bBuild )
13671428 case 7 : m_fnDist = IPBinaryFloatDistance<true , BinaryDotProduct16<true >>; break ;
13681429 }
13691430
1370- if ( !bBuild )
1371- {
1372- if ( bUseSSE )
1373- m_eDistFuncId = bNeedResiduals ? DistFuncId_e::IP_BINARY_SIMD16_RESIDUALS : DistFuncId_e::IP_BINARY_SIMD16 ;
1374- else
1375- m_eDistFuncId = DistFuncId_e::IP_BINARY_GENERIC ;
1376- }
1431+ if ( bUseSSE )
1432+ m_eDistFuncId = bNeedResiduals ? DistFuncId_e::IP_BINARY_SIMD16_RESIDUALS : DistFuncId_e::IP_BINARY_SIMD16 ;
1433+ else
1434+ m_eDistFuncId = DistFuncId_e::IP_BINARY_GENERIC ;
13771435#endif
13781436}
13791437
@@ -1391,13 +1449,10 @@ L2SpaceBinaryFloat_c::L2SpaceBinaryFloat_c ( size_t uDim, bool bBuild )
13911449 , m_tDistFuncParam ( uDim )
13921450{
13931451#if defined(USE_SIMDE)
1394- if ( bBuild )
1395- m_fnDist = L2BinaryFloatDistance<true ,BinaryDotProduct>;
1396- else
1397- {
1398- m_fnDist = L2BinaryFloatDistance<false ,BinaryDotProduct>;
1399- m_eDistFuncId = DistFuncId_e::L2_BINARY_GENERIC ;
1400- }
1452+ m_fnDist = bBuild
1453+ ? L2BinaryFloatDistance<true ,BinaryDotProduct>
1454+ : L2BinaryFloatDistance<false ,BinaryDotProduct>;
1455+ m_eDistFuncId = DistFuncId_e::L2_BINARY_GENERIC ;
14011456#else
14021457 int iBytes = ( uDim+7 ) >> 3 ;
14031458 bool bUseSSE = iBytes>=16 ;
@@ -1417,13 +1472,10 @@ L2SpaceBinaryFloat_c::L2SpaceBinaryFloat_c ( size_t uDim, bool bBuild )
14171472 case 7 : m_fnDist = L2BinaryFloatDistance<true , BinaryDotProduct16<true >>; break ;
14181473 }
14191474
1420- if ( !bBuild )
1421- {
1422- if ( bUseSSE )
1423- m_eDistFuncId = bNeedResiduals ? DistFuncId_e::L2_BINARY_SIMD16_RESIDUALS : DistFuncId_e::L2_BINARY_SIMD16 ;
1424- else
1425- m_eDistFuncId = DistFuncId_e::L2_BINARY_GENERIC ;
1426- }
1475+ if ( bUseSSE )
1476+ m_eDistFuncId = bNeedResiduals ? DistFuncId_e::L2_BINARY_SIMD16_RESIDUALS : DistFuncId_e::L2_BINARY_SIMD16 ;
1477+ else
1478+ m_eDistFuncId = DistFuncId_e::L2_BINARY_GENERIC ;
14271479#endif
14281480}
14291481
0 commit comments