// scnn_demo.mbt -- Minimum-viable SCNN training demo (v0.21.0).
//
// End-to-end verification that the project pipeline works:
// 1. Build synthetic 28x28 dataset (10 classes, fixed patterns).
// 2. Forward through SimpleCNN.
// 3. Compute softmax + cross-entropy loss.
// 4. Backward through the chain (returns per-layer gradients).
// 5. Verify loss is finite / reasonable / positive (sanity).
// 6. Verify grads are finite (no NaN / Inf).
///|
/// Build a 10-class synthetic 28x28 dataset. Each class has a fixed
/// binary pattern plus ~5% pixel-flip noise.
pub fn synthetic_mnist_build() -> (Array[Array[Float]], Array[Int]) {
let n_per_class = 10
let classes = 10
let total = n_per_class * classes
let images : Array[Array[Float]] = Array::make(total, [])
let labels : Array[Int] = Array::make(total, 0)
let rng = Xoshiro::new(42UL)
let mut cls = 0
while cls < classes {
let mut sample = 0
while sample < n_per_class {
let idx = cls * n_per_class + sample
labels[idx] = cls
let img : Array[Float] = Array::make(28 * 28, 0.0F)
if cls == 0 {
let mut c = 0
while c < 28 {
img[14 * 28 + c] = 1.0F
c = c + 1
}
} else if cls == 1 {
let mut r = 0
while r < 28 {
img[r * 28 + 14] = 1.0F
r = r + 1
}
} else if cls == 2 {
let mut r = 0
while r < 14 {
let mut c = 0
while c < 28 {
img[r * 28 + c] = 1.0F
c = c + 1
}
r = r + 1
}
} else if cls == 3 {
let mut r = 14
while r < 28 {
let mut c = 0
while c < 28 {
img[r * 28 + c] = 1.0F
c = c + 1
}
r = r + 1
}
} else if cls == 4 {
let mut r = 9
while r < 19 {
let mut c = 9
while c < 19 {
img[r * 28 + c] = 1.0F
c = c + 1
}
r = r + 1
}
} else if cls == 5 {
let mut i = 0
while i < 28 {
img[i * 28 + i] = 1.0F
i = i + 1
}
} else if cls == 6 {
let mut i = 0
while i < 28 {
img[i * 28 + (27 - i)] = 1.0F
i = i + 1
}
} else if cls == 7 {
let mut r = 0
while r < 28 {
img[r * 28 + 14] = 1.0F
r = r + 1
}
let mut c = 0
while c < 28 {
img[14 * 28 + c] = 1.0F
c = c + 1
}
} else if cls == 8 {
let mut r = 0
while r < 28 {
let mut c = 0
while c < 28 {
let is_border = r < 2 || r > 25 || c < 2 || c > 25
if is_border { img[r * 28 + c] = 1.0F }
c = c + 1
}
r = r + 1
}
} else {
let mut i = 0
while i < 28 {
img[i * 28 + i] = 1.0F
img[i * 28 + (27 - i)] = 1.0F
i = i + 1
}
}
// Add noise.
let mut i = 0
while i < 28 * 28 {
let u = next_f32(rng)
let original = img[i]
let is_bright = original > 0.5F
let flipped : Float = if is_bright { 0.0F } else { 1.0F }
let chosen : Float = if u < 0.05F { flipped } else { original }
img[i] = chosen
i = i + 1
}
images[idx] = img
sample = sample + 1
}
cls = cls + 1
}
(images, labels)
}
///|
/// One-hot encode an Int label.
pub fn one_hot(label : Int, num_classes : Int) -> Array[Float] {
let row : Array[Float] = Array::make(num_classes, 0.0F)
row[label] = 1.0F
row
}
///|
/// Numerically-stable softmax.
pub fn softmax_row(x : Array[Float]) -> Array[Float] {
let n = x.length()
let out : Array[Float] = Array::make(n, 0.0F)
let mut max_v = x[0]
let mut i = 1
while i < n {
if x[i] > max_v { max_v = x[i] }
i = i + 1
}
let mut sum = 0.0F
let mut j = 0
while j < n {
out[j] = expf(x[j] - max_v)
sum = sum + out[j]
j = j + 1
}
let mut k = 0
while k < n {
out[k] = out[k] / sum
k = k + 1
}
out
}
///|
/// Cross-entropy for one row: loss = -sum(target * log(prob)).
pub fn cross_entropy_one(prob : Array[Float], target : Array[Float]) -> Float {
let mut s = 0.0F
let mut j = 0
while j < prob.length() {
s = s - target[j] * logf(prob[j])
j = j + 1
}
s
}
///|
/// Run a single forward + loss + backward pass. Returns
/// (loss, total_params_with_grads).
pub fn scnn_single_step() -> (Float, Int) {
let (images, labels) = synthetic_mnist_build()
let num_classes = 10
let batch = images.length()
let model = SimpleCNN::new(1, 8, 16, 7UL)
let flat_input : Array[Float] = Array::make(batch * 28 * 28, 0.0F)
let mut i = 0
while i < batch {
let img = images[i]
let mut j = 0
while j < 28 * 28 {
flat_input[i * 28 * 28 + j] = img[j]
j = j + 1
}
i = i + 1
}
let (logits_flat, caches, _, _, _, _) = simple_cnn_forward(
model, flat_input, batch,
)
let mut total_loss = 0.0F
let d_logits : Array[Float] = Array::make(batch * num_classes, 0.0F)
let mut ii = 0
while ii < batch {
let row : Array[Float] = Array::make(num_classes, 0.0F)
let mut jj = 0
while jj < num_classes {
row[jj] = logits_flat[ii * num_classes + jj]
jj = jj + 1
}
let probs = softmax_row(row)
let target = one_hot(labels[ii], num_classes)
total_loss = total_loss + cross_entropy_one(probs, target)
let mut k = 0
while k < num_classes {
d_logits[ii * num_classes + k] = probs[k] - target[k]
k = k + 1
}
ii = ii + 1
}
let mean_loss = total_loss / Float::from_int(batch)
let (_d_input, grads) = simple_cnn_backward(
model, caches, d_logits, batch,
)
let mut grad_count = 0
let mut g_idx = 0
while g_idx < grads.length() {
match grads[g_idx] {
Conv2d(d_w, _d_b) => {
let mut k = 0
while k < d_w.length() {
if d_w[k] != 0.0F { grad_count = grad_count + 1 }
k = k + 1
}
}
Linear(d_w, _d_b) => {
let mut k = 0
while k < d_w.length() {
if d_w[k] != 0.0F { grad_count = grad_count + 1 }
k = k + 1
}
}
BatchNorm2d(d_g, _d_b) => {
let mut k = 0
while k < d_g.length() {
if d_g[k] != 0.0F { grad_count = grad_count + 1 }
k = k + 1
}
}
_ => ()
}
g_idx = g_idx + 1
}
let _ = _d_input
(mean_loss, grad_count)
}
///|
/// Sanity check that the loss is finite, positive, and bounded.
pub fn scnn_loss_is_reasonable() -> Bool {
let (loss, _) = scnn_single_step()
if loss <= 0.0F { return false }
if loss != loss { return false }
if loss > 5.0F * 2.303F { return false }
true
}
///|
/// Check that running forward multiple times with the same input gives
/// identical outputs (determinism check).
pub fn scnn_deterministic_forward() -> Bool {
let model = SimpleCNN::new(1, 8, 16, 7UL)
let input : Array[Float] = Array::make(2 * 28 * 28, 0.5F)
let (out1, _, _, _, _, _) = simple_cnn_forward(model, input, 2)
let (out2, _, _, _, _, _) = simple_cnn_forward(model, input, 2)
if out1.length() != out2.length() { return false }
let mut i = 0
while i < out1.length() {
if out1[i] != out2[i] { return false }
i = i + 1
}
true
}
// ---- SCNN v2 (v0.27.1): full K-step training with accuracy ----
///|
/// Predict argmax class for a batch of logits (length batch * n_classes).
pub fn scnn_predict_batch(
logits_flat : Array[Float],
batch : Int,
n_classes : Int,
) -> Array[Int] {
let preds : Array[Float] = Array::make(batch, 0.0F)
let mut i = 0
while i < batch {
let mut best = 0
let mut best_val = logits_flat[i * n_classes + 0]
let mut c = 1
while c < n_classes {
let v = logits_flat[i * n_classes + c]
if v > best_val {
best_val = v
best = c
}
c = c + 1
}
preds[i] = Float::from_int(best)
i = i + 1
}
// Cast Float to Int via reading + casting
let out : Array[Int] = Array::make(batch, 0)
let mut j = 0
while j < batch {
out[j] = preds[j].to_int()
j = j + 1
}
out
}
///|
/// Compute classification accuracy on a batch.
pub fn scnn_accuracy(
logits_flat : Array[Float],
labels : Array[Int],
batch : Int,
n_classes : Int,
) -> Float {
let preds = scnn_predict_batch(logits_flat, batch, n_classes)
let mut correct = 0
let mut i = 0
while i < batch {
if preds[i] == labels[i] { correct = correct + 1 }
i = i + 1
}
Float::from_int(correct) / Float::from_int(batch)
}
///|
/// Run a single forward + loss + backward pass; apply SGD updates to
/// the model's layers in place. Returns (loss, accuracy) on the batch.
pub fn scnn_train_step_with_sgd(
model : SimpleCNN,
input : Array[Float],
labels : Array[Int],
batch : Int,
n_classes : Int,
lr : Float,
) -> (Float, Float) {
let (logits_flat, caches, _, _, _, _) = simple_cnn_forward(
model, input, batch,
)
// CE gradient (per-row: per-sample).
let d_logits : Array[Float] = Array::make(batch * n_classes, 0.0F)
let mut total_loss = 0.0F
let mut ii = 0
while ii < batch {
let row : Array[Float] = Array::make(n_classes, 0.0F)
let mut jj = 0
while jj < n_classes {
row[jj] = logits_flat[ii * n_classes + jj]
jj = jj + 1
}
let probs = softmax_row(row)
let target = one_hot(labels[ii], n_classes)
total_loss = total_loss + cross_entropy_one(probs, target)
let mut k = 0
while k < n_classes {
d_logits[ii * n_classes + k] = probs[k] - target[k]
k = k + 1
}
ii = ii + 1
}
let mean_loss = total_loss / Float::from_int(batch)
// Backward.
let (_d_input, grads) = simple_cnn_backward(model, caches, d_logits, batch)
ignore(_d_input)
// SGD update — walk per-layer params by matching on Layer enum
// variant (shadow binding gives access to the variant's payload).
let clip = 1.0F
let mut g_idx = 0
while g_idx < grads.length() {
let d_grad = grads[g_idx]
let d_layer = model.layers[g_idx]
match (d_layer, d_grad) {
(Conv2d(param), LayerGrad::Conv2d(d_w, d_b)) => {
clip_grad_array(d_w, clip)
clip_grad_array(d_b, clip)
for i in 0.. {
clip_grad_array(d_w, clip)
clip_grad_array(d_b, clip)
for i in 0.. {
clip_grad_array(d_g, clip)
clip_grad_array(d_b, clip)
for i in 0.. ()
}
g_idx = g_idx + 1
}
let acc = scnn_accuracy(logits_flat, labels, batch, n_classes)
(mean_loss, acc)
}
///|
/// Run K SGD training steps on the synthetic dataset. Returns
/// (initial_loss, final_loss, initial_acc, final_acc).
pub fn scnn_train_n_steps(
n_steps : Int,
lr : Float,
seed : UInt64,
) -> (Float, Float, Float, Float) {
let (images, labels) = synthetic_mnist_build()
let batch = images.length()
let n_classes = 10
let model = SimpleCNN::new(1, 8, 16, seed)
// Flatten images into a single batch.
let flat_input : Array[Float] = Array::make(batch * 28 * 28, 0.0F)
let mut i = 0
while i < batch {
let img = images[i]
let mut j = 0
while j < 28 * 28 {
flat_input[i * 28 * 28 + j] = img[j]
j = j + 1
}
i = i + 1
}
// Initial loss + accuracy (no updates).
let (init_loss, init_acc) = scnn_train_step_with_sgd(
model, flat_input, labels, batch, n_classes, 0.0F,
)
// K real steps.
for _k in 0..