Class ov::pass::RestoreTracedBatch#
-
class RestoreTracedBatch : public ov::pass::MatcherPass#
Restores the batch that tracing froze to one in a
window_reversesubgraph.A traced
window_reversemay compute its batch asint(windows.shape[0] / (H * W / ws / ws)), which tracing with batch one turns into a constant. Only this subgraph is matched: a leading constant one cannot be told apart from an intentional collapse in general, so the pass is not a generic batch restoration.In both pinned targets the baked constant one is replaced with the batch read from the
Parametershape, reusing theGather(ShapeOf(Parameter), 0)that the lastReshapealready takes it from, optionally through aConvert. The targets are edited in place: the baked dimension belongs to the shape tensor, so everyReshapereading it needs the same restored batch.Before
Every node between the two pinnedwindows [B * nW, ws, ws, C] | Reshape(Concat(1, H / ws, W / ws, ws, ws, -1)) (batch pinned by tracing) | Transpose(axis 0 preserved) | Reshape(Concat(1, H, W, -1)) (batch pinned by tracing) | Roll(non-leading axes) (optional, shifted windows) | Reshape(Concat(Gather(ShapeOf(Parameter), 0), H * W, C))
Reshapes and the last one must have a single consumer.After
windows [B * nW, ws, ws, C] | Reshape(Concat(Gather(ShapeOf(Parameter), 0), H / ws, W / ws, ws, ws, -1)) | Transpose(axis 0 preserved) | Reshape(Concat(Gather(ShapeOf(Parameter), 0), H, W, -1)) | Roll(non-leading axes) | Reshape(Concat(Gather(ShapeOf(Parameter), 0), H * W, C))