| 76 | } |
| 77 | |
| 78 | void PatternSolver::resolvePatterns(luci::Module *module) |
| 79 | { |
| 80 | _frozen._node_to_param.clear(); |
| 81 | |
| 82 | for (auto pattern : _options._patterns) |
| 83 | { |
| 84 | std::unique_ptr<pattern::PatternResolver> resolver; |
| 85 | switch (pattern) |
| 86 | { |
| 87 | case QuantizationPattern::Q8LayerNormWithQ16Variance: |
| 88 | resolver = std::make_unique<pattern::Q8LayerNormWithQ16VarianceResolver>(); |
| 89 | break; |
| 90 | case QuantizationPattern::Q8SoftmaxWithQ16SubExp: |
| 91 | resolver = std::make_unique<pattern::Q8SoftmaxWithQ16SubExpResolver>(); |
| 92 | break; |
| 93 | default: |
| 94 | throw std::runtime_error("Unsupported pattern to resolve"); |
| 95 | } |
| 96 | |
| 97 | auto const resolved = resolver->resolve(module); |
| 98 | for (const auto &node_param : resolved) |
| 99 | { |
| 100 | auto const frozen = _frozen._node_to_param.find(node_param.first); |
| 101 | if (frozen == _frozen._node_to_param.end()) |
| 102 | { |
| 103 | // node was not previously defined - just set it (no ambiguity) |
| 104 | _frozen._node_to_param[node_param.first] = node_param.second; |
| 105 | } |
| 106 | else if (frozen->second.dtype != node_param.second.dtype) |
| 107 | { |
| 108 | // ambiguity (incoming description conflicts with current) |
| 109 | throw std::runtime_error("Resolved patterns contradict each other"); |
| 110 | } |
| 111 | } |
| 112 | } |
| 113 | } |