1717package org .apache .lucene .sandbox .codecs .faiss ;
1818
1919import static java .lang .foreign .ValueLayout .ADDRESS ;
20+ import static java .lang .foreign .ValueLayout .JAVA_BYTE ;
2021import static java .lang .foreign .ValueLayout .JAVA_FLOAT ;
2122import static java .lang .foreign .ValueLayout .JAVA_INT ;
2223import static java .lang .foreign .ValueLayout .JAVA_LONG ;
3233import java .lang .invoke .MethodHandles ;
3334import java .lang .invoke .MethodType ;
3435import java .nio .ByteOrder ;
35- import java .nio .FloatBuffer ;
36- import java .nio .LongBuffer ;
3736import java .util .Arrays ;
3837import java .util .Locale ;
3938import org .apache .lucene .index .FloatVectorValues ;
@@ -221,16 +220,22 @@ public static MemorySegment createIndex(
221220
222221 // Allocate docs in native memory
223222 MemorySegment docs = temp .allocate (JAVA_FLOAT , (long ) size * dimension );
224- FloatBuffer docsBuffer = docs .asByteBuffer ().order (ByteOrder .nativeOrder ()).asFloatBuffer ();
223+ long docsOffset = 0 ;
224+ long perDocByteSize = dimension * JAVA_FLOAT .byteSize ();
225225
226226 // Allocate ids in native memory
227227 MemorySegment ids = temp .allocate (JAVA_LONG , size );
228- LongBuffer idsBuffer = ids . asByteBuffer (). order ( ByteOrder . nativeOrder ()). asLongBuffer () ;
228+ int idsIndex = 0 ;
229229
230230 KnnVectorValues .DocIndexIterator iterator = floatVectorValues .iterator ();
231231 for (int i = iterator .nextDoc (); i != NO_MORE_DOCS ; i = iterator .nextDoc ()) {
232- idsBuffer .put (oldToNewDocId .apply (i ));
233- docsBuffer .put (floatVectorValues .vectorValue (iterator .index ()));
232+ int id = oldToNewDocId .apply (i );
233+ ids .setAtIndex (JAVA_LONG , idsIndex , id );
234+ idsIndex ++;
235+
236+ float [] vector = floatVectorValues .vectorValue (iterator .index ());
237+ MemorySegment .copy (vector , 0 , docs , JAVA_FLOAT , docsOffset , vector .length );
238+ docsOffset += perDocByteSize ;
234239 }
235240
236241 // Train index
@@ -254,18 +259,12 @@ private static long writeBytes(
254259 inputPointer = inputPointer .reinterpret (size );
255260
256261 if (size <= BUFFER_SIZE ) { // simple case, avoid buffering
257- byte [] bytes = new byte [(int ) size ];
258- inputPointer .asSlice (0 , size ).asByteBuffer ().order (ByteOrder .nativeOrder ()).get (bytes );
259- output .writeBytes (bytes , bytes .length );
262+ output .writeBytes (inputPointer .toArray (JAVA_BYTE ), (int ) size );
260263 } else { // copy buffered number of bytes repeatedly
261264 byte [] bytes = new byte [BUFFER_SIZE ];
262265 for (long offset = 0 ; offset < size ; offset += BUFFER_SIZE ) {
263266 int length = (int ) Math .min (size - offset , BUFFER_SIZE );
264- inputPointer
265- .asSlice (offset , length )
266- .asByteBuffer ()
267- .order (ByteOrder .nativeOrder ())
268- .get (bytes , 0 , length );
267+ MemorySegment .copy (inputPointer , JAVA_BYTE , offset , bytes , 0 , length );
269268 output .writeBytes (bytes , length );
270269 }
271270 }
@@ -282,21 +281,13 @@ private static long readBytes(
282281 if (size <= BUFFER_SIZE ) { // simple case, avoid buffering
283282 byte [] bytes = new byte [(int ) size ];
284283 input .readBytes (bytes , 0 , bytes .length );
285- outputPointer
286- .asSlice (0 , bytes .length )
287- .asByteBuffer ()
288- .order (ByteOrder .nativeOrder ())
289- .put (bytes );
284+ MemorySegment .copy (bytes , 0 , outputPointer , JAVA_BYTE , 0 , bytes .length );
290285 } else { // copy buffered number of bytes repeatedly
291286 byte [] bytes = new byte [BUFFER_SIZE ];
292287 for (long offset = 0 ; offset < size ; offset += BUFFER_SIZE ) {
293288 int length = (int ) Math .min (size - offset , BUFFER_SIZE );
294289 input .readBytes (bytes , 0 , length );
295- outputPointer
296- .asSlice (offset , length )
297- .asByteBuffer ()
298- .order (ByteOrder .nativeOrder ())
299- .put (bytes , 0 , length );
290+ MemorySegment .copy (bytes , 0 , outputPointer , JAVA_BYTE , offset , length );
300291 }
301292 }
302293 return numItems ;
@@ -411,8 +402,7 @@ public static void indexSearch(
411402 };
412403
413404 // Allocate queries in native memory
414- MemorySegment queries = temp .allocate (JAVA_FLOAT , query .length );
415- queries .asByteBuffer ().order (ByteOrder .nativeOrder ()).asFloatBuffer ().put (query );
405+ MemorySegment queries = temp .allocateFrom (JAVA_FLOAT , query );
416406
417407 // Faiss knn search
418408 int k = knnCollector .k ();
@@ -427,10 +417,9 @@ public static void indexSearch(
427417 MemorySegment pointer = temp .allocate (ADDRESS );
428418
429419 long [] bits = fixedBitSet .getBits ();
430- MemorySegment nativeBits = temp .allocate (JAVA_LONG , bits .length );
431-
432- // Use LITTLE_ENDIAN to convert long[] -> uint8_t*
433- nativeBits .asByteBuffer ().order (ByteOrder .LITTLE_ENDIAN ).asLongBuffer ().put (bits );
420+ MemorySegment nativeBits =
421+ // Use LITTLE_ENDIAN to convert long[] -> uint8_t*
422+ temp .allocateFrom (JAVA_LONG .withOrder (ByteOrder .LITTLE_ENDIAN ), bits );
434423
435424 callAndHandleError (ID_SELECTOR_BITMAP_NEW , pointer , fixedBitSet .length (), nativeBits );
436425 MemorySegment idSelectorBitmapPointer =
@@ -458,13 +447,9 @@ public static void indexSearch(
458447 idsPointer );
459448 }
460449
461- // Retrieve scores
462- float [] distances = new float [k ];
463- distancesPointer .asByteBuffer ().order (ByteOrder .nativeOrder ()).asFloatBuffer ().get (distances );
464-
465- // Retrieve ids
466- long [] ids = new long [k ];
467- idsPointer .asByteBuffer ().order (ByteOrder .nativeOrder ()).asLongBuffer ().get (ids );
450+ // Retrieve scores and ids
451+ float [] distances = distancesPointer .toArray (JAVA_FLOAT );
452+ long [] ids = idsPointer .toArray (JAVA_LONG );
468453
469454 // Record hits
470455 for (int i = 0 ; i < k ; i ++) {
0 commit comments