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 @@ -4,8 +4,8 @@ import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertFailsWith
import kotlin.test.assertNotNull
import kotlin.test.Ignore
import sk.ainet.context.DirectCpuExecutionContext
import sk.ainet.context.Phase
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.types.FP32
import sk.ainet.lang.tensor.Tensor
Expand Down Expand Up @@ -37,33 +37,32 @@ class BatchNormalizationTest {
}
}

@Ignore
@Test
fun train_then_eval_works_and_preserves_shape() {
val exec = DirectCpuExecutionContext()
val x = makeInput2x2(exec)
val trainExec = DirectCpuExecutionContext(phase = Phase.TRAIN)
val evalExec = DirectCpuExecutionContext(phase = Phase.EVAL)
val x = makeInput2x2(trainExec)
val bn = BatchNormalization<FP32, Float>(
numFeatures = 2,
affine = false,
name = "bn"
)
// training pass initializes running stats
bn.train()
val yTrain = bn.forward(x, exec)
val yTrain = bn.forward(x, trainExec)
assertNotNull(yTrain)
assertEquals(x.shape, yTrain.shape)

// eval should now work using running stats
bn.eval()
val yEval = bn.forward(x, exec)
val yEval = bn.forward(x, evalExec)
assertNotNull(yEval)
assertEquals(x.shape, yEval.shape)
}

@Ignore
@Test
fun simple_2x2_batch_is_normalized_per_channel() {
val exec = DirectCpuExecutionContext()
val exec = DirectCpuExecutionContext(phase = Phase.TRAIN)
val x = makeInput2x2(exec)
val bn = BatchNormalization<FP32, Float>(
numFeatures = 2,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ import kotlin.test.Test
import kotlin.test.Ignore
import kotlin.test.assertTrue

@Ignore
@Ignore("GraphExecution DSL tests are parked until the API drift is resolved")
class GraphExecutionDSLTest {
@Test
fun placeholder() {
Expand Down Expand Up @@ -298,4 +298,4 @@ class GraphExecutionDSLTest {

}
*/
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ import kotlin.test.assertTrue
* val graph = model.toGraph()
* graph.toGraphviz()
*/
@Ignore
@Ignore("Graphviz snippet is parked until the MnistMpl graph API drift is resolved")
class MnistMplGraphvizTest {


Expand All @@ -30,4 +30,4 @@ class MnistMplGraphvizTest {
// Keeping body minimal to allow compilation when @Ignore handling differs across targets.
assertTrue(true)
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ import kotlin.test.Test
import kotlin.test.Ignore
import kotlin.test.assertTrue

@Ignore
@Ignore("Placeholder suite parked until the JVM-specific execution helper is migrated")
class TapeToGraphUnitTests {
@Test
fun placeholder() {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ import kotlin.test.Ignore

class OnnxResourceReadTest {

@Ignore
@Ignore("Requires run14.onnx test fixture, which is not checked into the repository")
@Test
fun `read run14 onnx from resources and build graph view`() {
val inputStream: InputStream = requireNotNull(javaClass.getResourceAsStream("/run14.onnx")) {
Expand All @@ -38,7 +38,7 @@ class OnnxResourceReadTest {
)
}

@Ignore
@Ignore("Requires run14.onnx test fixture, which is not checked into the repository")
@Test
fun `run14 onnx ops are covered by importer mapping`() {
val bytes = loadResourceBytes("run14.onnx")
Expand Down
Loading