| 403 | } |
| 404 | |
| 405 | void 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, |
no test coverage detected