Skip to content

Commit 6cb820d

Browse files
authored
Merge pull request #57 from fsprojects/daily-test-improver-utils-mnist
Daily Test Coverage Improver: Research and Documentation
2 parents 9395380 + 7e75a38 commit 6cb820d

2 files changed

Lines changed: 155 additions & 0 deletions

File tree

tests/Furnace.Tests/TestData.fs

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
namespace Tests
77

8+
open System
89
open System.IO
910
open System.IO.Compression
1011
open System.Text
@@ -56,6 +57,10 @@ type TestData () =
5657
Assert.AreEqual(classesCorrect, classes)
5758
Assert.AreEqual(classNamesCorrect, classNames)
5859

60+
// Note: Removed problematic MNIST constructor tests that required network access
61+
// The MNIST constructor immediately downloads data, so unit testing its properties
62+
// is not feasible without network access or modifying the implementation
63+
5964
[<Test>]
6065
member _.TestCIFAR10Dataset () =
6166
// Note: this test can fail if https://www.cs.toronto.edu/~kriz website goes down or file urls change
Lines changed: 150 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,150 @@
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.0f; 2.0f; 3.0f; 4.0f|]
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.0f, extracted.[0])
28+
Assert.AreEqual(4.0f, 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 = [|255uy; 128uy; 64uy; 0uy|]
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(255uy, extracted.[0])
80+
Assert.AreEqual(0uy, extracted.[3])
81+
82+
[<Test>]
83+
member _.TestGetTypedValuesInt8() =
84+
// Test GetTypedValues extension method on int8 tensors
85+
let values = [|127y; -128y; 0y; 42y|]
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(127y, extracted.[0])
93+
Assert.AreEqual(42y, extracted.[3])
94+
95+
[<Test>]
96+
member _.TestGetTypedValuesInt16() =
97+
// Test GetTypedValues extension method on int16 tensors
98+
let values = [|1000s; -1000s; 0s; 32000s|]
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(1000s, extracted.[0])
106+
Assert.AreEqual(32000s, extracted.[3])
107+
108+
[<Test>]
109+
member _.TestGetTypedValuesInt64() =
110+
// Test GetTypedValues extension method on int64 tensors
111+
let values = [|1000000L; -1000000L; 0L; 9223372036854775807L|]
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(1000000L, extracted.[0])
119+
Assert.AreEqual(9223372036854775807L, extracted.[3])
120+
121+
[<Test>]
122+
member _.TestGetTypedValuesWrongType() =
123+
// Test that casting to wrong type throws exception
124+
let values = [|1.0f; 2.0f; 3.0f; 4.0f|]
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.0f|]
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.0f, 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

Comments
 (0)