(self, net, is_test, out_blob=None)
| 66 | optimizer=model.NoOptim) |
| 67 | |
| 68 | def _add_ops(self, net, is_test, out_blob=None): |
| 69 | original_input_blob = self.input_record.field_blobs() |
| 70 | input_blob = net.NextScopedBlob('expand_input') |
| 71 | if len(self.input_shape) == 1: |
| 72 | input_blob = net.ExpandDims(original_input_blob, |
| 73 | dims=[2, 3]) |
| 74 | else: |
| 75 | input_blob = original_input_blob[0] |
| 76 | |
| 77 | if out_blob is None: |
| 78 | bn_output = self.output_schema.field_blobs() |
| 79 | else: |
| 80 | bn_output = out_blob |
| 81 | if is_test: |
| 82 | output_blobs = bn_output |
| 83 | else: |
| 84 | output_blobs = bn_output + [self.rm, self.riv, |
| 85 | net.NextScopedBlob('bn_saved_mean'), |
| 86 | net.NextScopedBlob('bn_saved_iv')] |
| 87 | |
| 88 | net.SpatialBN([input_blob, self.scale, |
| 89 | self.bias, self.rm, self.riv], |
| 90 | output_blobs, |
| 91 | momentum=self.momentum, |
| 92 | is_test=is_test, |
| 93 | order=self.order) |
| 94 | |
| 95 | if len(self.input_shape) == 1: |
| 96 | net.Squeeze(bn_output, |
| 97 | bn_output, |
| 98 | dims=[2, 3]) |
| 99 | |
| 100 | def add_train_ops(self, net): |
| 101 | self._add_ops(net, is_test=False) |
no test coverage detected