| 188 | } |
| 189 | |
| 190 | TVector<THolder<IColumnPrinter>> InitializeColumnWriter( |
| 191 | const TEvalColumnsInfo& evalColumnsInfo, |
| 192 | NPar::ILocalExecutor* executor, |
| 193 | const TVector<TVector<TString>>& outputColumns, // [modelIdx] |
| 194 | const TDataProvider& pool, |
| 195 | TIntrusivePtr<IPoolColumnsPrinter> poolColumnsPrinter, |
| 196 | std::pair<int, int> testFileWhichOf, |
| 197 | ui64 docIdOffset, |
| 198 | bool* needColumnsPrinterPtr, |
| 199 | TMaybe<std::pair<size_t, size_t>> evalParameters, |
| 200 | double binClassLogitThreshold) { |
| 201 | |
| 202 | TFeatureIdToDesc featureIdToDesc = GetFeatureIdToDesc(pool); |
| 203 | |
| 204 | TVector<THolder<IColumnPrinter>> columnPrinter; |
| 205 | |
| 206 | const auto targetDim = pool.RawTargetData.GetTargetDimension(); |
| 207 | const bool isMultiTarget = targetDim > 1; |
| 208 | const auto modelCount = evalColumnsInfo.Approxes.size(); |
| 209 | CB_ENSURE_INTERNAL( |
| 210 | modelCount == outputColumns.size() |
| 211 | && modelCount == evalColumnsInfo.LossFunctions.size() |
| 212 | && modelCount == evalColumnsInfo.LabelHelpers.size(), |
| 213 | "Invalid evalColumnsInfo" |
| 214 | ); |
| 215 | |
| 216 | *needColumnsPrinterPtr = false; |
| 217 | |
| 218 | for (auto modelIdx : xrange(modelCount)) { |
| 219 | const auto& lossFunction = evalColumnsInfo.LossFunctions[modelIdx]; |
| 220 | const bool isMultiLabel = !lossFunction.empty() && IsMultiLabelObjective(lossFunction); |
| 221 | TMaybe<TString> modelName; |
| 222 | if (modelCount > 1) { |
| 223 | modelName = TString("Model") + ToString(modelIdx); |
| 224 | } |
| 225 | const auto& approx = evalColumnsInfo.Approxes[modelIdx]; |
| 226 | const auto& labelHelper = evalColumnsInfo.LabelHelpers[modelIdx]; |
| 227 | for (const auto& outputColumn : outputColumns[modelIdx]) { |
| 228 | EPredictionType type; |
| 229 | if (TryFromString<EPredictionType>(outputColumn, type)) { |
| 230 | PushBackEvalPrinters( |
| 231 | approx.GetRawValuesConstRef(), |
| 232 | type, |
| 233 | lossFunction, |
| 234 | modelName, |
| 235 | isMultiTarget, |
| 236 | approx.GetEnsemblesCount(), |
| 237 | labelHelper, |
| 238 | evalParameters, |
| 239 | &columnPrinter, |
| 240 | executor, |
| 241 | binClassLogitThreshold |
| 242 | ); |
| 243 | continue; |
| 244 | } |
| 245 | EColumn outputType; |
| 246 | if (TryFromString<EColumn>(ToCanonicalColumnName(outputColumn), outputType)) { |
| 247 | if (outputType == EColumn::Label) { |
no test coverage detected