| 744 | } |
| 745 | |
| 746 | std::vector<float> classify_embedding( |
| 747 | const EcapaWeights & weights, |
| 748 | const std::vector<float> & normalized_classifier_weight, |
| 749 | std::vector<float> & centered, |
| 750 | const std::vector<float> & embedding) { |
| 751 | const auto start = Clock::now(); |
| 752 | if (weights.embedding_global_mean.size() != static_cast<size_t>(kEmbeddingDim)) { |
| 753 | throw std::runtime_error("embedding global mean is missing from checkpoint"); |
| 754 | } |
| 755 | if (normalized_classifier_weight.empty() || normalized_classifier_weight.size() % static_cast<size_t>(kEmbeddingDim) != 0) { |
| 756 | throw std::runtime_error("normalized classifier weights are missing"); |
| 757 | } |
| 758 | const int64_t class_count = static_cast<int64_t>(normalized_classifier_weight.size() / static_cast<size_t>(kEmbeddingDim)); |
| 759 | if (centered.size() != static_cast<size_t>(kEmbeddingDim)) { |
| 760 | centered.resize(static_cast<size_t>(kEmbeddingDim)); |
| 761 | } |
| 762 | float emb_norm_sq = 0.0f; |
| 763 | for (int64_t i = 0; i < kEmbeddingDim; ++i) { |
| 764 | centered[static_cast<size_t>(i)] = embedding[static_cast<size_t>(i)] - weights.embedding_global_mean[static_cast<size_t>(i)]; |
| 765 | emb_norm_sq += centered[static_cast<size_t>(i)] * centered[static_cast<size_t>(i)]; |
| 766 | } |
| 767 | const float inv_emb_norm = 1.0f / std::sqrt(std::max(emb_norm_sq, kStatsEps)); |
| 768 | std::vector<float> logits(static_cast<size_t>(class_count), 0.0f); |
| 769 | for (int64_t cls = 0; cls < class_count; ++cls) { |
| 770 | const float * row = normalized_classifier_weight.data() + static_cast<size_t>(cls * kEmbeddingDim); |
| 771 | float dot = 0.0f; |
| 772 | for (int64_t i = 0; i < kEmbeddingDim; ++i) { |
| 773 | dot += centered[static_cast<size_t>(i)] * row[static_cast<size_t>(i)]; |
| 774 | } |
| 775 | logits[static_cast<size_t>(cls)] = dot * inv_emb_norm; |
| 776 | } |
| 777 | debug::timing_log_scalar("ecapa.classifier_ms", engine::debug::elapsed_ms(start, Clock::now())); |
| 778 | return logits; |
| 779 | } |
| 780 | |
| 781 | } // namespace |
| 782 |
no test coverage detected