MCPcopy Create free account
hub / github.com/catboost/catboost / SumModelsParams

Function SumModelsParams

catboost/libs/model/model.cpp:1730–1834  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1728}
1729
1730static void SumModelsParams(
1731 const TVector<const TFullModel*> modelVector,
1732 THashMap<TString, TString>* modelInfo
1733) {
1734 TMaybe<TString> classParams;
1735
1736 auto dimensionsCount = modelVector.back()->GetDimensionsCount();
1737
1738 for (auto modelIdx : xrange(modelVector.size())) {
1739 const auto& modelInfo = modelVector[modelIdx]->ModelInfo;
1740 bool paramFound = false;
1741 for (const auto& paramName : {"class_params", "multiclass_params"}) {
1742 if (modelInfo.contains(paramName)) {
1743 if (classParams) {
1744 CB_ENSURE(
1745 modelInfo.at(paramName) == *classParams,
1746 "Cannot sum models with different class params"
1747 );
1748 } else if ((modelIdx == 0) || (dimensionsCount == 1)) {
1749 // it is ok only for 1-dimensional models to have only some classParams specified
1750 classParams = modelInfo.at(paramName);
1751 } else {
1752 CB_ENSURE(false, "Cannot sum multidimensional models with and without class params");
1753 }
1754 paramFound = true;
1755 break;
1756 }
1757 }
1758 if ((modelIdx != 0) && classParams && !paramFound && (dimensionsCount > 1)) {
1759 CB_ENSURE(false, "Cannot sum multidimensional models with and without class params");
1760 }
1761 }
1762
1763 if (classParams) {
1764 (*modelInfo)["class_params"] = *classParams;
1765 } else {
1766 /* One-dimensional models.
1767 * If class labels for binary classification are present they must be the same
1768 */
1769
1770 TMaybe<TVector<NJson::TJsonValue>> sumClassLabels;
1771
1772 for (const TFullModel* model : modelVector) {
1773 TVector<NJson::TJsonValue> classLabels = model->GetModelClassLabels();
1774 if (classLabels) {
1775 CB_ENSURE(classLabels.size() == 2, "Expect exactly two class labels in binary classification");
1776
1777 if (sumClassLabels) {
1778 CB_ENSURE(classLabels == *sumClassLabels, "Cannot sum models with different class labels");
1779 } else {
1780 sumClassLabels = std::move(classLabels);
1781 }
1782 }
1783 }
1784
1785 if (sumClassLabels) {
1786 TString& paramsString = (*modelInfo)["params"];
1787 NJson::TJsonValue paramsJson;

Callers 1

SumModelsFunction · 0.85

Calls 15

xrangeFunction · 0.85
ReadTJsonValueFunction · 0.85
GetLossDescriptionFunction · 0.85
GetModelClassLabelsMethod · 0.80
AppendValueMethod · 0.80
moveFunction · 0.50
ToStringFunction · 0.50
AllOfFunction · 0.50
GetDimensionsCountMethod · 0.45
backMethod · 0.45
sizeMethod · 0.45
containsMethod · 0.45

Tested by

no test coverage detected