Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package sk.ainet.io.gguf
import sk.ainet.io.RandomAccessSource
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.storage.*
import sk.ainet.lang.types.*

/**
* Streaming GGUF reader that parses metadata without loading the entire file.
Expand Down Expand Up @@ -118,7 +119,7 @@ public class StreamingGGUFReader private constructor(
val shape = Shape(*tensor.shape.map { it.toInt() }.toIntArray())
return TensorStorage(
shape = shape,
logicalType = ggmlTypeToLogical(tensor.tensorType),
dtype = ggmlTypeToDType(tensor.tensorType),
encoding = ggmlTypeToEncoding(tensor.tensorType, tensor.nBytes),
buffer = BufferHandle.Borrowed(bytes, isMutable = false),
placement = Placement.CPU_HEAP
Expand Down Expand Up @@ -148,7 +149,7 @@ public class StreamingGGUFReader private constructor(
val shape = Shape(*tensor.shape.map { it.toInt() }.toIntArray())
return TensorStorage(
shape = shape,
logicalType = ggmlTypeToLogical(tensor.tensorType),
dtype = ggmlTypeToDType(tensor.tensorType),
encoding = ggmlTypeToEncoding(tensor.tensorType, tensor.nBytes),
buffer = BufferHandle.FileBacked(
path = filePath,
Expand All @@ -159,17 +160,17 @@ public class StreamingGGUFReader private constructor(
)
}

private fun ggmlTypeToLogical(type: GGMLQuantizationType): LogicalDType = when (type) {
GGMLQuantizationType.F32 -> LogicalDType.FLOAT32
GGMLQuantizationType.F16 -> LogicalDType.FLOAT16
GGMLQuantizationType.BF16 -> LogicalDType.BFLOAT16
GGMLQuantizationType.F64 -> LogicalDType.FLOAT64
GGMLQuantizationType.I8 -> LogicalDType.INT8
GGMLQuantizationType.I16 -> LogicalDType.INT16
GGMLQuantizationType.I32 -> LogicalDType.INT32
GGMLQuantizationType.I64 -> LogicalDType.INT64
private fun ggmlTypeToDType(type: GGMLQuantizationType): DType = when (type) {
GGMLQuantizationType.F32 -> FP32
GGMLQuantizationType.F16 -> FP16
GGMLQuantizationType.BF16 -> BF16
GGMLQuantizationType.F64 -> FP64
GGMLQuantizationType.I8 -> Int8
GGMLQuantizationType.I16 -> Int16
GGMLQuantizationType.I32 -> Int32
GGMLQuantizationType.I64 -> Int64
// Quantized types logically represent floats
else -> LogicalDType.FLOAT32
else -> FP32
}

private fun ggmlTypeToEncoding(type: GGMLQuantizationType, nBytes: Long): TensorEncoding = when (type) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,14 @@ package sk.ainet.io.gguf
import org.junit.Test
import sk.ainet.io.JvmFileBackedResolver
import sk.ainet.io.JvmRandomAccessSource
import sk.ainet.lang.types.FP32
import sk.ainet.lang.tensor.storage.*
import java.io.File
import java.io.RandomAccessFile
import java.nio.ByteBuffer
import java.nio.ByteOrder
import kotlin.test.assertEquals
import kotlin.test.assertSame
import kotlin.test.assertFalse
import kotlin.test.assertTrue

Expand Down Expand Up @@ -106,6 +108,7 @@ class StorageIntegrationTest {
// F32 tensor
val f32Storage = reader.loadTensorStorage("weight_f32")
assertEquals(LogicalDType.FLOAT32, f32Storage.logicalType)
assertSame(FP32, f32Storage.dtype)
assertEquals(TensorEncoding.Dense(4), f32Storage.encoding)
assertEquals(Ownership.BORROWED, f32Storage.ownership)
assertEquals(16L, f32Storage.physicalBytes)
Expand All @@ -114,7 +117,8 @@ class StorageIntegrationTest {

// Q8_0 tensor
val q80Storage = reader.loadTensorStorage("weight_q80")
assertEquals(LogicalDType.FLOAT32, q80Storage.logicalType)
assertEquals(LogicalDType.FLOAT32, q80Storage.logicalType) // packed weights are logically FP32
assertSame(FP32, q80Storage.dtype)
assertEquals(TensorEncoding.Q8_0, q80Storage.encoding)
assertEquals(Ownership.BORROWED, q80Storage.ownership)
assertEquals(34L, q80Storage.physicalBytes)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import sk.ainet.io.RandomAccessSource
import sk.ainet.io.model.DataType
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.storage.*
import sk.ainet.lang.types.*

/**
* Streaming SafeTensors reader that parses metadata without loading tensor data.
Expand Down Expand Up @@ -106,7 +107,7 @@ public class StreamingSafeTensorsReader private constructor(
val shape = Shape(*tensor.shape.map { it.toInt() }.toIntArray())
return TensorStorage(
shape = shape,
logicalType = safeTensorsTypeToLogical(tensor.dataType),
dtype = safeTensorsTypeToDType(tensor.dataType),
encoding = safeTensorsTypeToEncoding(tensor.dataType),
buffer = BufferHandle.Borrowed(bytes, isMutable = false),
placement = Placement.CPU_HEAP
Expand All @@ -133,7 +134,7 @@ public class StreamingSafeTensorsReader private constructor(
val shape = Shape(*tensor.shape.map { it.toInt() }.toIntArray())
return TensorStorage(
shape = shape,
logicalType = safeTensorsTypeToLogical(tensor.dataType),
dtype = safeTensorsTypeToDType(tensor.dataType),
encoding = safeTensorsTypeToEncoding(tensor.dataType),
buffer = BufferHandle.FileBacked(
path = filePath,
Expand All @@ -144,21 +145,21 @@ public class StreamingSafeTensorsReader private constructor(
)
}

private fun safeTensorsTypeToLogical(type: DataType): LogicalDType = when (type) {
DataType.FLOAT32 -> LogicalDType.FLOAT32
DataType.FLOAT64 -> LogicalDType.FLOAT64
DataType.FLOAT16 -> LogicalDType.FLOAT16
DataType.BFLOAT16 -> LogicalDType.BFLOAT16
DataType.INT8 -> LogicalDType.INT8
DataType.INT16 -> LogicalDType.INT16
DataType.INT32 -> LogicalDType.INT32
DataType.INT64 -> LogicalDType.INT64
DataType.UINT8 -> LogicalDType.UINT8
DataType.UINT16 -> LogicalDType.UINT16
DataType.UINT32 -> LogicalDType.UINT32
DataType.UINT64 -> LogicalDType.UINT64
DataType.BOOL -> LogicalDType.UINT8
else -> LogicalDType.INT8 // fallback for UNKNOWN
private fun safeTensorsTypeToDType(type: DataType): DType = when (type) {
DataType.FLOAT32 -> FP32
DataType.FLOAT64 -> FP64
DataType.FLOAT16 -> FP16
DataType.BFLOAT16 -> BF16
DataType.INT8 -> Int8
DataType.INT16 -> Int16
DataType.INT32 -> Int32
DataType.INT64 -> Int64
DataType.UINT8 -> UInt8
DataType.UINT16 -> UInt16
DataType.UINT32 -> UInt32
DataType.UINT64 -> UInt64
DataType.BOOL -> UInt8
else -> Int8 // fallback for UNKNOWN
}

private fun safeTensorsTypeToEncoding(type: DataType): TensorEncoding = when (type) {
Expand Down
Loading
Loading