call method
Implementation
Tensor call(Tensor x, int t) {
if (x.shape.length != 4 || x.shape[1] != 1) {
throw ArgumentError('TinyUNet: expected [N, 1, H, W]; got ${x.shape}');
}
if (t < 0 || t >= totalTimesteps) {
throw ArgumentError('TinyUNet: t=$t out of range [0, $totalTimesteps)');
}
final n = x.shape[0];
var h = down1(x).relu();
h = down2(h).relu();
h = mid(h);
// Time embedding: broadcast a `[N, 2*hidden]` vector across the
// spatial axes of the mid feature map.
final tScalar = t / totalTimesteps;
final tIn = Tensor.fromList(
[n, 1],
List<double>.filled(n, tScalar),
device: x.device,
);
final tEmb = timeProj(tIn); // [N, 2*hidden]
final ch = tEmb.shape[1];
final sh = h.shape[2];
final sw = h.shape[3];
// Broadcast add: replicate tEmb across [H, W] on host, then re-lift.
final tHost = Tensor.noGrad(() => tEmb.toList());
final broadcast = Float32List(n * ch * sh * sw);
for (int ni = 0; ni < n; ni++) {
for (int c = 0; c < ch; c++) {
final val = tHost[ni * ch + c];
final base = ((ni * ch + c) * sh) * sw;
for (int i = 0; i < sh * sw; i++) {
broadcast[base + i] = val;
}
}
}
final tEmbBroadcast = Tensor.fromFloat32List(
[n, ch, sh, sw],
broadcast,
device: x.device,
);
h = (h + tEmbBroadcast).relu();
h = up1(h).relu();
return up2(h);
}