Skip to content

Commit 24c557b

Browse files
committed
Changed ArrayToBuffer to require dimensions if len(flatValues) > 1.
1 parent 01748f8 commit 24c557b

2 files changed

Lines changed: 50 additions & 47 deletions

File tree

pkg/pjrt/buffers.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -282,8 +282,8 @@ func ScalarToBufferOnDeviceNum[T dtypes.Supported](client *Client, deviceNum int
282282
// It is a shortcut to Client.BufferFromHost call with default parameters.
283283
// If you need more control where the value will be used you'll have to use Client.BufferFromHost instead.
284284
func ArrayToBuffer[T dtypes.Supported](client *Client, flatValues []T, dimensions ...int) (b *Buffer, err error) {
285-
if len(dimensions) == 0 {
286-
dimensions = []int{len(flatValues)}
285+
if len(dimensions) == 0 && len(flatValues) != 1 {
286+
return nil, errors.Errorf("ArrayToBuffer not given any dimensions (indicating a scalar), but len(flatValues) == %d", len(flatValues))
287287
}
288288
return client.BufferFromHost().FromFlatDataWithDimensions(flatValues, dimensions).Done()
289289
}

pkg/pjrt/buffers_test.go

Lines changed: 48 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -22,56 +22,59 @@ func TestScalarDataToRaw(t *testing.T) {
2222
func testTransfersImpl[T interface {
2323
float64 | float32 | int64 | int8
2424
}](t *testing.T, client *Client) {
25-
// Transfer arrays.
26-
input := []T{1, 2, 3}
27-
fmt.Printf("From %#v\n", input)
28-
buffer, err := ArrayToBuffer(client, input, 3, 1)
29-
requireNoError(t, err)
30-
assertFalse(t, buffer.IsShared())
25+
var baseT T
26+
t.Run(fmt.Sprintf("%T", baseT), func(t *testing.T) {
27+
// Transfer arrays.
28+
input := []T{1, 2, 3}
29+
fmt.Printf("From %#v\n", input)
30+
buffer, err := ArrayToBuffer(client, input, 3, 1)
31+
requireNoError(t, err)
32+
assertFalse(t, buffer.IsShared())
3133

32-
output, outputDims, err := BufferToArray[T](buffer)
33-
requireNoError(t, err)
34-
fmt.Printf("\t> output=%#v\n", output)
35-
assertEqualSlice(t, input, output)
36-
assertEqualSlice(t, []int{3, 1}, outputDims)
34+
output, outputDims, err := BufferToArray[T](buffer)
35+
requireNoError(t, err)
36+
fmt.Printf("\t> output=%#v\n", output)
37+
assertEqualSlice(t, input, output)
38+
assertEqualSlice(t, []int{3, 1}, outputDims)
3739

38-
flat, outputDims, err := buffer.ToFlatDataAndDimensions()
39-
requireNoError(t, err)
40-
assertEqualSlice(t, input, flat.([]T))
41-
assertEqualSlice(t, []int{3, 1}, outputDims)
40+
flat, outputDims, err := buffer.ToFlatDataAndDimensions()
41+
requireNoError(t, err)
42+
assertEqualSlice(t, input, flat.([]T))
43+
assertEqualSlice(t, []int{3, 1}, outputDims)
4244

43-
gotDevice, err := buffer.Device()
44-
requireNoError(t, err)
45-
wantDevice := client.AddressableDevices()[0]
46-
assertEqual(t, wantDevice.LocalHardwareID(), gotDevice.LocalHardwareID())
47-
assertEqual(t, 0, client.NumForDevice(gotDevice))
48-
49-
// Try an invalid transfer: it should complain about the invalid dtype.
50-
_, _, err = BufferToArray[complex128](buffer)
51-
fmt.Printf("\t> expected wrong dtype error: %v\n", err)
52-
requireError(t, err)
53-
54-
// Transfer scalars.
55-
from := T(13)
56-
fmt.Printf("From %T(%v)\n", from, from)
57-
buffer, err = ScalarToBuffer(client, from)
58-
requireNoError(t, err)
59-
to, err := BufferToScalar[T](buffer)
60-
requireNoError(t, err)
61-
fmt.Printf("\t> got %v\n", to)
62-
assertEqual(t, from, to)
45+
gotDevice, err := buffer.Device()
46+
requireNoError(t, err)
47+
wantDevice := client.AddressableDevices()[0]
48+
assertEqual(t, wantDevice.LocalHardwareID(), gotDevice.LocalHardwareID())
49+
assertEqual(t, 0, client.NumForDevice(gotDevice))
50+
51+
// Try an invalid transfer: it should complain about the invalid dtype.
52+
_, _, err = BufferToArray[complex128](buffer)
53+
fmt.Printf("\t> expected wrong dtype error: %v\n", err)
54+
requireError(t, err)
55+
56+
// Transfer scalars.
57+
from := T(13)
58+
fmt.Printf("From %T(%v)\n", from, from)
59+
buffer, err = ScalarToBuffer(client, from)
60+
requireNoError(t, err)
61+
to, err := BufferToScalar[T](buffer)
62+
requireNoError(t, err)
63+
fmt.Printf("\t> got %v\n", to)
64+
assertEqual(t, from, to)
6365

64-
// ArrayToBuffer can also be used to transfer a scalar.
65-
from = T(19)
66-
fmt.Printf("From %T(%v)\n", from, from)
67-
buffer, err = ArrayToBuffer(client, []T{from})
68-
requireNoError(t, err)
66+
// ArrayToBuffer can also be used to transfer a scalar.
67+
from = T(19)
68+
fmt.Printf("From %T(%v)\n", from, from)
69+
buffer, err = ArrayToBuffer(client, []T{from})
70+
requireNoError(t, err)
6971

70-
flatValues, dimensions, err := BufferToArray[T](buffer) // Check that it actually returns a scalar.
71-
requireNoError(t, err)
72-
assertLen(t, dimensions, 0) // That means, it is a scalar.
73-
fmt.Printf("\t> got %v\n", flatValues[0])
74-
assertEqual(t, from, flatValues[0])
72+
flatValues, dimensions, err := BufferToArray[T](buffer) // Check that it actually returns a scalar.
73+
requireNoError(t, err)
74+
assertLen(t, dimensions, 0) // That means, it is a scalar.
75+
fmt.Printf("\t> got %v\n", flatValues[0])
76+
assertEqual(t, from, flatValues[0])
77+
})
7578
}
7679

7780
func TestTransfers(t *testing.T) {

0 commit comments

Comments
 (0)