| 50 | } // namespace |
| 51 | |
| 52 | ::tensorflow::Status ResolveConstantPack::Run(Model* model, |
| 53 | std::size_t op_index, |
| 54 | bool* modified) { |
| 55 | *modified = false; |
| 56 | auto it = model->operators.begin() + op_index; |
| 57 | const auto* base_op = it->get(); |
| 58 | if (base_op->type != OperatorType::kPack) { |
| 59 | return ::tensorflow::Status::OK(); |
| 60 | } |
| 61 | const auto* op = static_cast<const PackOperator*>(base_op); |
| 62 | |
| 63 | CHECK_GE(op->inputs.size(), 1); |
| 64 | CHECK_EQ(op->outputs.size(), 1); |
| 65 | auto& output_array = model->GetArray(op->outputs[0]); |
| 66 | if (output_array.data_type == ArrayDataType::kNone) { |
| 67 | // Yield until the output type has been set by PropagateArrayDataTypes |
| 68 | return ::tensorflow::Status::OK(); |
| 69 | } |
| 70 | |
| 71 | if (!output_array.has_shape()) { |
| 72 | // Yield until the output shape has been set by PropagateFixedShapes |
| 73 | return ::tensorflow::Status::OK(); |
| 74 | } |
| 75 | |
| 76 | for (const auto& input : op->inputs) { |
| 77 | if (!IsConstantParameterArray(*model, input)) { |
| 78 | // Yield if any input is mutable |
| 79 | return ::tensorflow::Status::OK(); |
| 80 | } |
| 81 | } |
| 82 | |
| 83 | int axis = op->axis; |
| 84 | if (axis < 0) { |
| 85 | // Handle negative axis |
| 86 | axis += model->GetArray(op->inputs[0]).shape().dims().size(); |
| 87 | } |
| 88 | CHECK_EQ(axis, 0) << "Packing only supported along 0th axis"; |
| 89 | |
| 90 | CHECK(!output_array.buffer); |
| 91 | switch (output_array.data_type) { |
| 92 | case ArrayDataType::kFloat: |
| 93 | Pack<ArrayDataType::kFloat>(model, *op); |
| 94 | break; |
| 95 | case ArrayDataType::kUint8: |
| 96 | Pack<ArrayDataType::kUint8>(model, *op); |
| 97 | break; |
| 98 | case ArrayDataType::kInt32: |
| 99 | Pack<ArrayDataType::kInt32>(model, *op); |
| 100 | break; |
| 101 | case ArrayDataType::kInt64: |
| 102 | Pack<ArrayDataType::kInt64>(model, *op); |
| 103 | break; |
| 104 | case ArrayDataType::kComplex64: |
| 105 | Pack<ArrayDataType::kComplex64>(model, *op); |
| 106 | break; |
| 107 | default: |
| 108 | LOG(FATAL) << "Unsupported data type given to Pack op with output \"" |
| 109 | << op->outputs[0] << "\""; |
nothing calls this directly
no test coverage detected