| 350 | |
| 351 | |
| 352 | class Resnet: |
| 353 | def run(self, img): |
| 354 | """ |
| 355 | 执行 ResNet 模型的前向传播。 |
| 356 | |
| 357 | 参数: |
| 358 | img (numpy.ndarray): 预处理后的输入图像。 |
| 359 | |
| 360 | 返回: |
| 361 | numpy.ndarray: 模型的输出。 |
| 362 | """ |
| 363 | # 初始卷积层、批归一化和ReLU层 |
| 364 | out = ComputeConvLayer(img, "conv1") |
| 365 | out = ComputeBatchNormLayer(out, "bn1") |
| 366 | out = ComputeReluLayer(out) |
| 367 | out = ComputeMaxPoolLayer(out) |
| 368 | |
| 369 | # 通过四个残差层(每层有多个残差块)处理数据 |
| 370 | # layer1 |
| 371 | out = ComputeBottleNeck(out, "layer1_bottleneck0", down_sample=True) |
| 372 | out = ComputeBottleNeck(out, "layer1_bottleneck1", down_sample=False) |
| 373 | out = ComputeBottleNeck(out, "layer1_bottleneck2", down_sample=False) |
| 374 | |
| 375 | # layer2 |
| 376 | out = ComputeBottleNeck(out, "layer2_bottleneck0", down_sample=True) |
| 377 | out = ComputeBottleNeck(out, "layer2_bottleneck1", down_sample=False) |
| 378 | out = ComputeBottleNeck(out, "layer2_bottleneck2", down_sample=False) |
| 379 | out = ComputeBottleNeck(out, "layer2_bottleneck3", down_sample=False) |
| 380 | |
| 381 | # layer3 |
| 382 | out = ComputeBottleNeck(out, "layer3_bottleneck0", down_sample=True) |
| 383 | out = ComputeBottleNeck(out, "layer3_bottleneck1", down_sample=False) |
| 384 | out = ComputeBottleNeck(out, "layer3_bottleneck2", down_sample=False) |
| 385 | out = ComputeBottleNeck(out, "layer3_bottleneck3", down_sample=False) |
| 386 | out = ComputeBottleNeck(out, "layer3_bottleneck4", down_sample=False) |
| 387 | out = ComputeBottleNeck(out, "layer3_bottleneck5", down_sample=False) |
| 388 | |
| 389 | # layer4 |
| 390 | out = ComputeBottleNeck(out, "layer4_bottleneck0", down_sample=True) |
| 391 | out = ComputeBottleNeck(out, "layer4_bottleneck1", down_sample=False) |
| 392 | out = ComputeBottleNeck(out, "layer4_bottleneck2", down_sample=False) |
| 393 | |
| 394 | # 平均池化和全连接层 |
| 395 | out = ComputeAvgPoolLayer(out) |
| 396 | out = ComputeFcLayer(out, "fc") |
| 397 | return out |
| 398 | |
| 399 | |
| 400 | # 导入时间模块,用来计算模型推理时间 |