| 488 | } |
| 489 | |
| 490 | std::vector<Float> getTopicPrior(const std::string& metadata, |
| 491 | const std::vector<std::string>& mdVec, |
| 492 | bool raw = false |
| 493 | ) const override |
| 494 | { |
| 495 | Vid xid = metadataDict.toWid(metadata); |
| 496 | if (xid == (Vid)-1) throw exc::InvalidArgument("unknown metadata '" + metadata + "'"); |
| 497 | |
| 498 | Vector xs = Vector::Zero(mdVecSize); |
| 499 | xs[0] = 1; |
| 500 | for (auto& m : mdVec) |
| 501 | { |
| 502 | Vid x = multiMetadataDict.toWid(m); |
| 503 | if (x == (Vid)-1) throw exc::InvalidArgument("unknown multi_metadata '" + m + "'"); |
| 504 | xs[x + 1] = 1; |
| 505 | } |
| 506 | |
| 507 | std::vector<Float> ret(this->K); |
| 508 | Eigen::Map<Vector> map{ ret.data(), (Eigen::Index)ret.size() }; |
| 509 | |
| 510 | if (raw) |
| 511 | { |
| 512 | map = lambda.middleCols(xid * mdVecSize, mdVecSize) * xs; |
| 513 | } |
| 514 | else |
| 515 | { |
| 516 | map = (lambda.middleCols(xid * mdVecSize, mdVecSize) * xs).array().exp() + alphaEps; |
| 517 | } |
| 518 | return ret; |
| 519 | } |
| 520 | |
| 521 | const Dictionary& getMetadataDict() const override { return metadataDict; } |
| 522 | const Dictionary& getMultiMetadataDict() const override { return multiMetadataDict; } |
nothing calls this directly
no test coverage detected