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

Function PostProcessingIndependent

catboost/libs/fstr/independent_tree_shap.cpp:405–488  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

403}
404
405void PostProcessingIndependent(
406 const TIndependentTreeShapParams& independentTreeShapParams,
407 const TVector<TVector<TVector<double>>>& shapValuesInternalForAllReferences,
408 const TVector<TVector<int>>& combinationClassFeatures,
409 size_t approxDimension,
410 size_t flatFeatureCount,
411 size_t documentIdx,
412 bool calcInternalValues,
413 const TVector<double>& bias,
414 TVector<TVector<double>>* shapValues
415) {
416 const size_t featureCount = calcInternalValues ? combinationClassFeatures.size() : flatFeatureCount;
417 const size_t referenceCount = independentTreeShapParams.ReferenceLeafIndicesForAllTrees[0].size();
418 EExplainableModelOutput modelOutputType = independentTreeShapParams.ModelOutputType;
419 const bool isNotRawOutputType = (EExplainableModelOutput::Raw != modelOutputType);
420 const auto& metric = *independentTreeShapParams.Metric.Get();
421 const auto& approxOfDataset = independentTreeShapParams.ApproxOfDataset;
422 const auto& approxOfReferenceDataset = independentTreeShapParams.ApproxOfReferenceDataset;
423 const auto& targetOfDataset = independentTreeShapParams.TargetOfDataset;
424 const auto& transformedTargetOfDataset = independentTreeShapParams.TransformedTargetOfDataset;
425 const bool isExplainMultiClassProbabilities = approxDimension > 1 && EExplainableModelOutput::Probability == modelOutputType;
426 TVector<TVector<double>> meanValuesProbabitiesForAllReference;
427 if (isExplainMultiClassProbabilities) {
428 meanValuesProbabitiesForAllReference = GetProbabilityMeanValues(shapValuesInternalForAllReferences, bias);
429 }
430 // prepare shap values for all references
431 TVector<TVector<TVector<double>>> shapValuesForAllReferences(approxDimension);
432 for (size_t dimension = 0; dimension < approxDimension; ++dimension) {
433 shapValuesForAllReferences[dimension].resize(referenceCount);
434 for (size_t referenceIdx = 0; referenceIdx < referenceCount; ++referenceIdx) {
435 shapValuesForAllReferences[dimension][referenceIdx] = calcInternalValues ?
436 shapValuesInternalForAllReferences[referenceIdx][dimension] :
437 GetUnpackedShapValues(
438 shapValuesInternalForAllReferences[referenceIdx][dimension],
439 combinationClassFeatures,
440 flatFeatureCount
441 );
442 }
443 }
444
445 const auto& probabilitiesOfReferenceDataset = independentTreeShapParams.ProbabilitiesOfReferenceDataset;
446 const bool isMultiTarget = (targetOfDataset.size() > 1);
447 for (size_t dimension = 0; dimension < approxDimension; ++dimension) {
448 TConstArrayRef<double> approxOfReferenceDatasetRef = MakeConstArrayRef(approxOfReferenceDataset[dimension]);
449 const double targetOfDocument = isMultiTarget ? targetOfDataset[dimension][documentIdx] : targetOfDataset[0][documentIdx];
450 const auto& transformedTargetOfReferenceDataset = isExplainMultiClassProbabilities ?
451 probabilitiesOfReferenceDataset[dimension] :
452 GetTransformData(
453 metric,
454 approxOfReferenceDatasetRef,
455 modelOutputType,
456 isNotRawOutputType,
457 targetOfDocument
458 );
459 const auto& meanValues = isExplainMultiClassProbabilities ?
460 meanValuesProbabitiesForAllReference[dimension] :
461 GetMeanValues(
462 metric,

Callers 1

Calls 9

GetProbabilityMeanValuesFunction · 0.85
GetUnpackedShapValuesFunction · 0.85
MakeConstArrayRefFunction · 0.85
GetTransformDataFunction · 0.85
GetMeanValuesFunction · 0.85
MakeArrayRefFunction · 0.85
sizeMethod · 0.45
GetMethod · 0.45
resizeMethod · 0.45

Tested by

no test coverage detected