| 530 | } |
| 531 | |
| 532 | Optional<BufferOffset<tflite::Buffer>> Translator::BuildBuffer( |
| 533 | Operation* inst) { |
| 534 | ElementsAttr attr; |
| 535 | if (auto cst = dyn_cast<mlir::ConstantOp>(inst)) { |
| 536 | // ConstantOp have ElementAttr at this point due to validation of the TFLite |
| 537 | // module. |
| 538 | attr = cst.getValue().cast<ElementsAttr>(); |
| 539 | } else if (auto cst = dyn_cast<mlir::TF::ConstOp>(inst)) { |
| 540 | attr = cst.value(); |
| 541 | } else if (auto cst = dyn_cast<tfl::ConstOp>(inst)) { |
| 542 | attr = cst.value(); |
| 543 | } else if (auto cst = dyn_cast<tfl::QConstOp>(inst)) { |
| 544 | attr = cst.value(); |
| 545 | } else if (auto cst = dyn_cast<tfl::SparseConstOp>(inst)) { |
| 546 | attr = cst.value(); |
| 547 | } else if (auto cst = dyn_cast<tfl::SparseQConstOp>(inst)) { |
| 548 | attr = cst.value(); |
| 549 | } else { |
| 550 | return empty_buffer_; |
| 551 | } |
| 552 | |
| 553 | tensorflow::Tensor tensor; |
| 554 | auto status = tensorflow::ConvertToTensor(attr, &tensor); |
| 555 | if (!status.ok()) { |
| 556 | inst->emitError( |
| 557 | Twine("failed to convert value attribute to tensor with error: " + |
| 558 | status.ToString())); |
| 559 | return llvm::None; |
| 560 | } |
| 561 | |
| 562 | // TensorFlow and TensorFlow Lite use different string encoding formats. |
| 563 | // Convert to TensorFlow Lite format is it's a constant string tensor. |
| 564 | if (tensor.dtype() == tensorflow::DT_STRING) { |
| 565 | ::tflite::DynamicBuffer dynamic_buffer; |
| 566 | auto flat = tensor.flat<::tensorflow::tstring>(); |
| 567 | for (int i = 0; i < flat.size(); ++i) { |
| 568 | const auto& str = flat(i); |
| 569 | dynamic_buffer.AddString(str.c_str(), str.length()); |
| 570 | } |
| 571 | char* tensor_buffer; |
| 572 | int bytes = dynamic_buffer.WriteToBuffer(&tensor_buffer); |
| 573 | auto buffer_data = |
| 574 | builder_.CreateVector(reinterpret_cast<uint8_t*>(tensor_buffer), bytes); |
| 575 | free(tensor_buffer); |
| 576 | return tflite::CreateBuffer(builder_, buffer_data); |
| 577 | } |
| 578 | |
| 579 | absl::string_view tensor_data = tensor.tensor_data(); |
| 580 | auto buffer_data = builder_.CreateVector( |
| 581 | reinterpret_cast<const uint8_t*>(tensor_data.data()), tensor_data.size()); |
| 582 | return tflite::CreateBuffer(builder_, buffer_data); |
| 583 | } |
| 584 | |
| 585 | Optional<BufferOffset<tflite::Tensor>> Translator::BuildTensor( |
| 586 | Value value, const std::string& name, unsigned buffer_idx) { |
nothing calls this directly
no test coverage detected