| 134 | } // namespace |
| 135 | |
| 136 | TEST(LayerNormPatternResolverTest, resolve_pattern) |
| 137 | { |
| 138 | auto m = luci::make_module(); |
| 139 | LayerNormTestGraph g; |
| 140 | g.init(); |
| 141 | g.transfer_to(m.get()); |
| 142 | |
| 143 | std::map<luci::CircleNode *, LayerParam> params; |
| 144 | mpqsolver::pattern::Q8LayerNormWithQ16VarianceResolver resolver; |
| 145 | EXPECT_NO_THROW({ params = resolver.resolve(m.get()); }); |
| 146 | |
| 147 | std::set<luci::CircleNode *> q16_nodes = {g.sub_squared, g.mean_as_variance, g.add_eps, g.rsqrt}; |
| 148 | std::set<luci::CircleNode *> q8_nodes = {g.mean_of_ifm, g.sub, g.mul}; |
| 149 | |
| 150 | // params of all valid layers are set |
| 151 | EXPECT_EQ(params.size(), q16_nodes.size() + q8_nodes.size()); |
| 152 | |
| 153 | for (auto param : params) |
| 154 | { |
| 155 | // params of all layers are set as prescribed |
| 156 | if (q16_nodes.find(param.first) != q16_nodes.end()) |
| 157 | { |
| 158 | EXPECT_STREQ(param.second.dtype.c_str(), "int16"); |
| 159 | } |
| 160 | else if (q8_nodes.find(param.first) != q8_nodes.end()) |
| 161 | { |
| 162 | EXPECT_STREQ(param.second.dtype.c_str(), "uint8"); |
| 163 | } |
| 164 | } |
| 165 | } |
| 166 | |
| 167 | TEST(LayerNormPatternResolverTest, resolve_pattern_NEG) |
| 168 | { |
nothing calls this directly
no test coverage detected