1+ // Copyright (c) 2016- University of Oxford (Atılım Güneş Baydin <gunes@robots.ox.ac.uk>)
2+ // and other contributors, see LICENSE in root of repository.
3+ //
4+ // BSD 2-Clause License. See LICENSE in root of repository.
5+
6+ namespace Tests
7+
8+ open System
9+ open Furnace
10+ open Furnace.Backends .Reference
11+ open NUnit.Framework
12+ open Tests.TestUtils
13+
14+ [<TestFixture>]
15+ type TestReferenceUtils () =
16+
17+ [<Test>]
18+ member _.TestGetTypedValuesFloat32 () =
19+ // Test GetTypedValues extension method on float32 tensors
20+ let values = [| 1.0 f; 2.0 f; 3.0 f; 4.0 f|]
21+ let shape = Shape.create [| 2 ; 2 |]
22+ let tensor = RawTensorFloat32( values, shape, Device.CPU) :> RawTensor
23+
24+ let extracted = tensor.GetTypedValues< float32>()
25+ Assert.AreEqual( values, extracted)
26+ Assert.AreEqual( 4 , extracted.Length)
27+ Assert.AreEqual( 1.0 f, extracted.[ 0 ])
28+ Assert.AreEqual( 4.0 f, extracted.[ 3 ])
29+
30+ [<Test>]
31+ member _.TestGetTypedValuesInt32 () =
32+ // Test GetTypedValues extension method on int32 tensors
33+ let values = [| 10 ; 20 ; 30 ; 40 |]
34+ let shape = Shape.create [| 2 ; 2 |]
35+ let tensor = RawTensorInt32( values, shape, Device.CPU) :> RawTensor
36+
37+ let extracted = tensor.GetTypedValues< int32>()
38+ Assert.AreEqual( values, extracted)
39+ Assert.AreEqual( 4 , extracted.Length)
40+ Assert.AreEqual( 10 , extracted.[ 0 ])
41+ Assert.AreEqual( 40 , extracted.[ 3 ])
42+
43+ [<Test>]
44+ member _.TestGetTypedValuesFloat64 () =
45+ // Test GetTypedValues extension method on float64 tensors
46+ let values = [| 1.5 ; 2.5 ; 3.5 ; 4.5 |]
47+ let shape = Shape.create [| 4 |]
48+ let tensor = RawTensorFloat64( values, shape, Device.CPU) :> RawTensor
49+
50+ let extracted = tensor.GetTypedValues< float64>()
51+ Assert.AreEqual( values, extracted)
52+ Assert.AreEqual( 4 , extracted.Length)
53+ Assert.AreEqual( 1.5 , extracted.[ 0 ])
54+ Assert.AreEqual( 4.5 , extracted.[ 3 ])
55+
56+ [<Test>]
57+ member _.TestGetTypedValuesBool () =
58+ // Test GetTypedValues extension method on bool tensors
59+ let values = [| true ; false ; true ; false |]
60+ let shape = Shape.create [| 2 ; 2 |]
61+ let tensor = RawTensorBool( values, shape, Device.CPU) :> RawTensor
62+
63+ let extracted = tensor.GetTypedValues< bool>()
64+ Assert.AreEqual( values, extracted)
65+ Assert.AreEqual( 4 , extracted.Length)
66+ Assert.AreEqual( true , extracted.[ 0 ])
67+ Assert.AreEqual( false , extracted.[ 3 ])
68+
69+ [<Test>]
70+ member _.TestGetTypedValuesByte () =
71+ // Test GetTypedValues extension method on byte tensors
72+ let values = [| 255 uy; 128 uy; 64 uy; 0 uy|]
73+ let shape = Shape.create [| 4 |]
74+ let tensor = RawTensorByte( values, shape, Device.CPU) :> RawTensor
75+
76+ let extracted = tensor.GetTypedValues< byte>()
77+ Assert.AreEqual( values, extracted)
78+ Assert.AreEqual( 4 , extracted.Length)
79+ Assert.AreEqual( 255 uy, extracted.[ 0 ])
80+ Assert.AreEqual( 0 uy, extracted.[ 3 ])
81+
82+ [<Test>]
83+ member _.TestGetTypedValuesInt8 () =
84+ // Test GetTypedValues extension method on int8 tensors
85+ let values = [| 127 y; - 128 y; 0 y; 42 y|]
86+ let shape = Shape.create [| 2 ; 2 |]
87+ let tensor = RawTensorInt8( values, shape, Device.CPU) :> RawTensor
88+
89+ let extracted = tensor.GetTypedValues< int8>()
90+ Assert.AreEqual( values, extracted)
91+ Assert.AreEqual( 4 , extracted.Length)
92+ Assert.AreEqual( 127 y, extracted.[ 0 ])
93+ Assert.AreEqual( 42 y, extracted.[ 3 ])
94+
95+ [<Test>]
96+ member _.TestGetTypedValuesInt16 () =
97+ // Test GetTypedValues extension method on int16 tensors
98+ let values = [| 1000 s; - 1000 s; 0 s; 32000 s|]
99+ let shape = Shape.create [| 4 |]
100+ let tensor = RawTensorInt16( values, shape, Device.CPU) :> RawTensor
101+
102+ let extracted = tensor.GetTypedValues< int16>()
103+ Assert.AreEqual( values, extracted)
104+ Assert.AreEqual( 4 , extracted.Length)
105+ Assert.AreEqual( 1000 s, extracted.[ 0 ])
106+ Assert.AreEqual( 32000 s, extracted.[ 3 ])
107+
108+ [<Test>]
109+ member _.TestGetTypedValuesInt64 () =
110+ // Test GetTypedValues extension method on int64 tensors
111+ let values = [| 1000000 L; - 1000000 L; 0 L; 9223372036854775807 L|]
112+ let shape = Shape.create [| 2 ; 2 |]
113+ let tensor = RawTensorInt64( values, shape, Device.CPU) :> RawTensor
114+
115+ let extracted = tensor.GetTypedValues< int64>()
116+ Assert.AreEqual( values, extracted)
117+ Assert.AreEqual( 4 , extracted.Length)
118+ Assert.AreEqual( 1000000 L, extracted.[ 0 ])
119+ Assert.AreEqual( 9223372036854775807 L, extracted.[ 3 ])
120+
121+ [<Test>]
122+ member _.TestGetTypedValuesWrongType () =
123+ // Test that casting to wrong type throws exception
124+ let values = [| 1.0 f; 2.0 f; 3.0 f; 4.0 f|]
125+ let shape = Shape.create [| 2 ; 2 |]
126+ let tensor = RawTensorFloat32( values, shape, Device.CPU) :> RawTensor
127+
128+ // This should throw when trying to cast float32 tensor to int32 values
129+ isException ( fun () -> tensor.GetTypedValues< int32>())
130+
131+ [<Test>]
132+ member _.TestGetTypedValuesScalar () =
133+ // Test GetTypedValues extension method on scalar tensors
134+ let values = [| 42.0 f|]
135+ let shape = Shape.create [||] // scalar shape
136+ let tensor = RawTensorFloat32( values, shape, Device.CPU) :> RawTensor
137+
138+ let extracted = tensor.GetTypedValues< float32>()
139+ Assert.AreEqual( 1 , extracted.Length)
140+ Assert.AreEqual( 42.0 f, extracted.[ 0 ])
141+
142+ [<Test>]
143+ member _.TestGetTypedValuesEmpty () =
144+ // Test GetTypedValues extension method on empty tensors
145+ let values = [||]
146+ let shape = Shape.create [| 0 |] // empty tensor
147+ let tensor = RawTensorFloat32( values, shape, Device.CPU) :> RawTensor
148+
149+ let extracted = tensor.GetTypedValues< float32>()
150+ Assert.AreEqual( 0 , extracted.Length)
0 commit comments