#include #include #include #include // 辅助函数:打印张量形状信息 void print_shape(const std::vector& shape) { std::cout << "["; for (size_t i = 0; i < shape.size(); ++i) { std::cout << shape[i]; if (i != shape.size() - 1) std::cout << ", "; } std::cout << "]" << std::endl; } // 核心函数:将 3x896x896 图像转换为 4096x2x3x14x14 Patches void image_to_patches( const float* input_data, // 输入图像指针 [3, 896, 896] float* output_data, // 输出 patches 指针 [4096, 2, 3, 14, 14] int img_c, int img_h, int img_w, // 输入图像维度: 3, 896, 896 int patch_h, int patch_w, // Patch 尺寸: 14, 14 int target_n, int target_t // 目标维度: 4096, 2 ) { // 计算每行的 patch 数量 int patches_per_row = img_w / patch_w; // 896 / 14 = 64 int patches_per_col = img_h / patch_h; // 896 / 14 = 64 // 校验目标 N 是否匹配 assert(patches_per_row * patches_per_col == target_n && "计算出的 Patch 总数与目标的 N=4096 不符!"); // 遍历每一个 Patch for (int p_idx = 0; p_idx < target_n; ++p_idx) { // 计算当前 patch 在原图上的起始坐标 int row = p_idx / patches_per_row; int col = p_idx % patches_per_row; int start_y = row * patch_h; int start_x = col * patch_w; // 遍历时间/副本维度 T (0 到 1) for (int t = 0; t < target_t; ++t) { // 遍历 RGB 三个通道 for (int c = 0; c < img_c; ++c) { // 遍历 Patch 内部像素 for (int ph = 0; ph < patch_h; ++ph) { for (int pw = 0; pw < patch_w; ++pw) { // --- 输入索引计算 [C, H, W] --- int in_y = start_y + ph; int in_x = start_x + pw; // 防止越界 (虽然这里整除刚好不会越界) if (in_y >= img_h || in_x >= img_w) continue; size_t input_idx = (size_t)c * img_h * img_w + (size_t)in_y * img_w + (size_t)in_x; // --- 输出索引计算 [N, T, C, h, w] --- size_t output_idx = 0; output_idx += (size_t)p_idx * target_t * img_c * patch_h * patch_w; output_idx += (size_t)t * img_c * patch_h * patch_w; output_idx += (size_t)c * patch_h * patch_w; output_idx += (size_t)ph * patch_w; output_idx += (size_t)pw; // 赋值 output_data[output_idx] = input_data[input_idx]; } } } } } } int main() { // 1. 定义输入和输出形状 const int IMG_C = 3; const int IMG_H = 896; const int IMG_W = 896; const int PATCH_H = 14; const int PATCH_W = 14; const int TARGET_N = 4096; // 64 * 64 const int TARGET_T = 2; // 2. 分配内存 // 输入大小: 3 * 896 * 896 = 2,408,448 std::vector input_image(IMG_C * IMG_H * IMG_W, 0.0f); // 输出大小: 4096 * 2 * 3 * 14 * 14 = 4,825,344 // 注意:输出大小是输入的两倍,因为 T=2 维度相当于存了两份所有的 patches std::vector output_patches(TARGET_N * TARGET_T * IMG_C * PATCH_H * PATCH_W, 0.0f); // 3. 填充一些测试数据 (例如给 R 通道赋值为 1.0) for (int i = 0; i < IMG_H * IMG_W; ++i) { input_image[i] = 1.0f; // R channel part (simplified indexing for init) } std::cout << "Input Image Shape: "; print_shape({IMG_C, IMG_H, IMG_W}); std::cout << "Input Element Count: " << input_image.size() << std::endl; // 4. 执行转换 image_to_patches( input_image.data(), output_patches.data(), IMG_C, IMG_H, IMG_W, PATCH_H, PATCH_W, TARGET_N, TARGET_T ); std::cout << "Output Patches Shape: "; print_shape({TARGET_N, TARGET_T, IMG_C, PATCH_H, PATCH_W}); std::cout << "Output Element Count: " << output_patches.size() << std::endl; // 5. 简单验证:检查第一个 patch 的第一个像素是否拷贝成功 // 根据我们的索引计算,output[0, 0, 0, 0, 0] 应该等于 input[0, 0, 0] float test_val = output_patches[0]; std::cout << "Verification: output_patches[0] = " << test_val << " (Expected 1.0)" << std::endl; return 0; }