import torch class PatchProjectionLayer(torch.nn.Module):     def __init__(self, patch_size, num_channels, embedding_dim):         super().__init__()         self.patch_size = patch_size         self.num_channels = num_channels         self.embedding_dim = embedding_dim         self.projection = torch.nn.Linear( patch_size * patch_size * num_channels, embedding_dim )     def forward(self, x):         batch_size, num_patches, channels, height, width = x.size()         x = x.view(batch_size, num_patches, -1)  # Flatten each patch         x = self.projection(x)  # Project each flattened patch         return x # Example Usage: batch_size = 1 num_patches = 9  # Total patches per image patch_size = 16  # 16x16 pixels per patch num_channels = 3  # RGB image embedding_dim = 768  # Size of the embedding vector projection_layer = PatchProjectionLayer(patch_size, num_channels, embedding_dim) patches = torch.rand( batch_size, num_patches, num_channels, patch_size, patch_size ) projected_embeddings = projection_layer(patches) print(projected_embeddings.shape) # This prints # torch.Size([1, 9, 768])