call method

Tensor call(
  1. Tensor x,
  2. int t
)

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);
}