@@ -22,56 +22,59 @@ func TestScalarDataToRaw(t *testing.T) {
2222func 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
7780func TestTransfers (t * testing.T ) {
0 commit comments