| 624 | return inputTensor; |
| 625 | } |
| 626 | bool Variable::input(VARP src) { |
| 627 | if (nullptr != mFrom->get()) { |
| 628 | MNN_ERROR("Can't input to no-input op\n"); |
| 629 | return false; |
| 630 | } |
| 631 | if (nullptr == src) { |
| 632 | /*Close the Input*/ |
| 633 | mFrom->visitOutputs([](EXPRP expr, int index) { |
| 634 | auto recurse = expr->mValid; expr->mValid = false; |
| 635 | return recurse; |
| 636 | }); |
| 637 | mFrom->mValid = false; |
| 638 | return false; |
| 639 | } |
| 640 | auto info = src->getInfo(); |
| 641 | std::shared_ptr<Variable::Info> tempInfo; |
| 642 | if (nullptr == info) { |
| 643 | tempInfo.reset(new Variable::Info); |
| 644 | tempInfo->size = 0; |
| 645 | tempInfo->type = halide_type_of<float>(); |
| 646 | info = tempInfo.get(); |
| 647 | } |
| 648 | auto dstInfo = getInfo(); |
| 649 | bool needChange = nullptr == dstInfo || info->order != dstInfo->order || info->dim.size() != dstInfo->dim.size() || info->type != dstInfo->type; |
| 650 | if (!needChange) { |
| 651 | for (int i=0; i<info->dim.size(); ++i) { |
| 652 | if (dstInfo->dim[i] != info->dim[i]) { |
| 653 | needChange = true; |
| 654 | break; |
| 655 | } |
| 656 | } |
| 657 | } |
| 658 | |
| 659 | if (!mFrom->mInside->mCache) { |
| 660 | ExecutorScope::Current()->makeCache({mFrom}, false); |
| 661 | } |
| 662 | if (needChange) { |
| 663 | mFrom->mInside->mOutputInfos[0] = *info; |
| 664 | Utils::releaseMemoryForHostTensor(mFrom->inside()->mOutputTensors[0]); |
| 665 | Utils::copyInfoToTensor(mFrom->inside()->mOutputTensors[0], mFrom->inside()->mOutputInfos.data()); |
| 666 | Utils::allocMemoryForHostTensor(mFrom->inside()->mOutputTensors[0]); |
| 667 | } |
| 668 | if (info->size) { |
| 669 | auto dstPtr = writeInternal(false); |
| 670 | auto srcPtr = src->readMap<void>(); |
| 671 | if (nullptr == dstPtr || nullptr == srcPtr) { |
| 672 | //MNN_ERROR("Alloc memory error or compute src error in Variable::Input\n"); |
| 673 | return false; |
| 674 | } |
| 675 | ::memcpy(dstPtr, srcPtr, info->size * info->type.bytes()); |
| 676 | } |
| 677 | if (needChange) { |
| 678 | mFrom->visitOutputs([](EXPRP expr, int index) { return expr->setInfoDirty(); }); |
| 679 | } else { |
| 680 | informDirty(); |
| 681 | } |
| 682 | mFrom->mInside->mContentDirty = false; |
| 683 | return true; |
no test coverage detected