mirror of
https://github.com/opencv/opencv.git
synced 2026-07-29 23:33:05 +04:00
add loading TensorFlow/Caffe net from memory buffer
add a corresponding test
This commit is contained in:
@@ -94,6 +94,17 @@ public:
|
||||
ReadNetParamsFromBinaryFileOrDie(caffeModel, &netBinary);
|
||||
}
|
||||
|
||||
CaffeImporter(const char *dataProto, size_t lenProto,
|
||||
const char *dataModel, size_t lenModel)
|
||||
{
|
||||
CV_TRACE_FUNCTION();
|
||||
|
||||
ReadNetParamsFromTextBufferOrDie(dataProto, lenProto, &net);
|
||||
|
||||
if (dataModel != NULL && lenModel > 0)
|
||||
ReadNetParamsFromBinaryBufferOrDie(dataModel, lenModel, &netBinary);
|
||||
}
|
||||
|
||||
void addParam(const Message &msg, const FieldDescriptor *field, cv::dnn::LayerParams ¶ms)
|
||||
{
|
||||
const Reflection *refl = msg.GetReflection();
|
||||
@@ -400,6 +411,15 @@ Net readNetFromCaffe(const String &prototxt, const String &caffeModel /*= String
|
||||
return net;
|
||||
}
|
||||
|
||||
Net readNetFromCaffe(const char *bufferProto, size_t lenProto,
|
||||
const char *bufferModel, size_t lenModel)
|
||||
{
|
||||
CaffeImporter caffeImporter(bufferProto, lenProto, bufferModel, lenModel);
|
||||
Net net;
|
||||
caffeImporter.populateNet(net);
|
||||
return net;
|
||||
}
|
||||
|
||||
#endif //HAVE_PROTOBUF
|
||||
|
||||
CV__DNN_EXPERIMENTAL_NS_END
|
||||
|
||||
@@ -1108,28 +1108,37 @@ const char* UpgradeV1LayerType(const V1LayerParameter_LayerType type) {
|
||||
|
||||
const int kProtoReadBytesLimit = INT_MAX; // Max size of 2 GB minus 1 byte.
|
||||
|
||||
bool ReadProtoFromBinary(ZeroCopyInputStream* input, Message *proto) {
|
||||
CodedInputStream coded_input(input);
|
||||
coded_input.SetTotalBytesLimit(kProtoReadBytesLimit, 536870912);
|
||||
|
||||
return proto->ParseFromCodedStream(&coded_input);
|
||||
}
|
||||
|
||||
bool ReadProtoFromTextFile(const char* filename, Message* proto) {
|
||||
std::ifstream fs(filename, std::ifstream::in);
|
||||
CHECK(fs.is_open()) << "Can't open \"" << filename << "\"";
|
||||
IstreamInputStream input(&fs);
|
||||
bool success = google::protobuf::TextFormat::Parse(&input, proto);
|
||||
fs.close();
|
||||
return success;
|
||||
return google::protobuf::TextFormat::Parse(&input, proto);
|
||||
}
|
||||
|
||||
bool ReadProtoFromBinaryFile(const char* filename, Message* proto) {
|
||||
std::ifstream fs(filename, std::ifstream::in | std::ifstream::binary);
|
||||
CHECK(fs.is_open()) << "Can't open \"" << filename << "\"";
|
||||
ZeroCopyInputStream* raw_input = new IstreamInputStream(&fs);
|
||||
CodedInputStream* coded_input = new CodedInputStream(raw_input);
|
||||
coded_input->SetTotalBytesLimit(kProtoReadBytesLimit, 536870912);
|
||||
IstreamInputStream raw_input(&fs);
|
||||
|
||||
bool success = proto->ParseFromCodedStream(coded_input);
|
||||
return ReadProtoFromBinary(&raw_input, proto);
|
||||
}
|
||||
|
||||
delete coded_input;
|
||||
delete raw_input;
|
||||
fs.close();
|
||||
return success;
|
||||
bool ReadProtoFromTextBuffer(const char* data, size_t len, Message* proto) {
|
||||
ArrayInputStream input(data, len);
|
||||
return google::protobuf::TextFormat::Parse(&input, proto);
|
||||
}
|
||||
|
||||
|
||||
bool ReadProtoFromBinaryBuffer(const char* data, size_t len, Message* proto) {
|
||||
ArrayInputStream raw_input(data, len);
|
||||
return ReadProtoFromBinary(&raw_input, proto);
|
||||
}
|
||||
|
||||
void ReadNetParamsFromTextFileOrDie(const char* param_file,
|
||||
@@ -1139,6 +1148,13 @@ void ReadNetParamsFromTextFileOrDie(const char* param_file,
|
||||
UpgradeNetAsNeeded(param_file, param);
|
||||
}
|
||||
|
||||
void ReadNetParamsFromTextBufferOrDie(const char* data, size_t len,
|
||||
NetParameter* param) {
|
||||
CHECK(ReadProtoFromTextBuffer(data, len, param))
|
||||
<< "Failed to parse NetParameter buffer";
|
||||
UpgradeNetAsNeeded("memory buffer", param);
|
||||
}
|
||||
|
||||
void ReadNetParamsFromBinaryFileOrDie(const char* param_file,
|
||||
NetParameter* param) {
|
||||
CHECK(ReadProtoFromBinaryFile(param_file, param))
|
||||
@@ -1146,6 +1162,13 @@ void ReadNetParamsFromBinaryFileOrDie(const char* param_file,
|
||||
UpgradeNetAsNeeded(param_file, param);
|
||||
}
|
||||
|
||||
void ReadNetParamsFromBinaryBufferOrDie(const char* data, size_t len,
|
||||
NetParameter* param) {
|
||||
CHECK(ReadProtoFromBinaryBuffer(data, len, param))
|
||||
<< "Failed to parse NetParameter buffer";
|
||||
UpgradeNetAsNeeded("memory buffer", param);
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
@@ -102,6 +102,18 @@ void ReadNetParamsFromTextFileOrDie(const char* param_file,
|
||||
void ReadNetParamsFromBinaryFileOrDie(const char* param_file,
|
||||
caffe::NetParameter* param);
|
||||
|
||||
// Read parameters from a memory buffer into a NetParammeter proto message.
|
||||
void ReadNetParamsFromBinaryBufferOrDie(const char* data, size_t len,
|
||||
caffe::NetParameter* param);
|
||||
void ReadNetParamsFromTextBufferOrDie(const char* data, size_t len,
|
||||
caffe::NetParameter* param);
|
||||
|
||||
// Utility functions used internally by Caffe and TensorFlow loaders
|
||||
bool ReadProtoFromTextFile(const char* filename, ::google::protobuf::Message* proto);
|
||||
bool ReadProtoFromBinaryFile(const char* filename, ::google::protobuf::Message* proto);
|
||||
bool ReadProtoFromTextBuffer(const char* data, size_t len, ::google::protobuf::Message* proto);
|
||||
bool ReadProtoFromBinaryBuffer(const char* data, size_t len, ::google::protobuf::Message* proto);
|
||||
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
@@ -449,6 +449,9 @@ void ExcludeLayer(tensorflow::GraphDef& net, const int layer_index, const int in
|
||||
class TFImporter : public Importer {
|
||||
public:
|
||||
TFImporter(const char *model, const char *config = NULL);
|
||||
TFImporter(const char *dataModel, size_t lenModel,
|
||||
const char *dataConfig = NULL, size_t lenConfig = 0);
|
||||
|
||||
void populateNet(Net dstNet);
|
||||
~TFImporter() {}
|
||||
|
||||
@@ -479,6 +482,15 @@ TFImporter::TFImporter(const char *model, const char *config)
|
||||
ReadTFNetParamsFromTextFileOrDie(config, &netTxt);
|
||||
}
|
||||
|
||||
TFImporter::TFImporter(const char *dataModel, size_t lenModel,
|
||||
const char *dataConfig, size_t lenConfig)
|
||||
{
|
||||
if (dataModel != NULL && lenModel > 0)
|
||||
ReadTFNetParamsFromBinaryBufferOrDie(dataModel, lenModel, &netBin);
|
||||
if (dataConfig != NULL && lenConfig > 0)
|
||||
ReadTFNetParamsFromTextBufferOrDie(dataConfig, lenConfig, &netTxt);
|
||||
}
|
||||
|
||||
void TFImporter::kernelFromTensor(const tensorflow::TensorProto &tensor, Mat &dstBlob)
|
||||
{
|
||||
MatShape shape;
|
||||
@@ -1326,5 +1338,14 @@ Net readNetFromTensorflow(const String &model, const String &config)
|
||||
return net;
|
||||
}
|
||||
|
||||
Net readNetFromTensorflow(const char* bufferModel, size_t lenModel,
|
||||
const char* bufferConfig, size_t lenConfig)
|
||||
{
|
||||
TFImporter importer(bufferModel, lenModel, bufferConfig, lenConfig);
|
||||
Net net;
|
||||
importer.populateNet(net);
|
||||
return net;
|
||||
}
|
||||
|
||||
CV__DNN_EXPERIMENTAL_NS_END
|
||||
}} // namespace
|
||||
|
||||
@@ -23,6 +23,7 @@ Implementation of various functions which are related to Tensorflow models readi
|
||||
|
||||
#include "graph.pb.h"
|
||||
#include "tf_io.hpp"
|
||||
#include "../caffe/caffe_io.hpp"
|
||||
#include "../caffe/glog_emulator.hpp"
|
||||
|
||||
namespace cv {
|
||||
@@ -36,41 +37,28 @@ using namespace ::google::protobuf::io;
|
||||
|
||||
const int kProtoReadBytesLimit = INT_MAX; // Max size of 2 GB minus 1 byte.
|
||||
|
||||
// TODO: remove Caffe duplicate
|
||||
bool ReadProtoFromBinaryFileTF(const char* filename, Message* proto) {
|
||||
std::ifstream fs(filename, std::ifstream::in | std::ifstream::binary);
|
||||
CHECK(fs.is_open()) << "Can't open \"" << filename << "\"";
|
||||
ZeroCopyInputStream* raw_input = new IstreamInputStream(&fs);
|
||||
CodedInputStream* coded_input = new CodedInputStream(raw_input);
|
||||
coded_input->SetTotalBytesLimit(kProtoReadBytesLimit, 536870912);
|
||||
|
||||
bool success = proto->ParseFromCodedStream(coded_input);
|
||||
|
||||
delete coded_input;
|
||||
delete raw_input;
|
||||
fs.close();
|
||||
return success;
|
||||
}
|
||||
|
||||
bool ReadProtoFromTextFileTF(const char* filename, Message* proto) {
|
||||
std::ifstream fs(filename, std::ifstream::in);
|
||||
CHECK(fs.is_open()) << "Can't open \"" << filename << "\"";
|
||||
IstreamInputStream input(&fs);
|
||||
bool success = google::protobuf::TextFormat::Parse(&input, proto);
|
||||
fs.close();
|
||||
return success;
|
||||
}
|
||||
|
||||
void ReadTFNetParamsFromBinaryFileOrDie(const char* param_file,
|
||||
tensorflow::GraphDef* param) {
|
||||
CHECK(ReadProtoFromBinaryFileTF(param_file, param))
|
||||
<< "Failed to parse GraphDef file: " << param_file;
|
||||
tensorflow::GraphDef* param) {
|
||||
CHECK(ReadProtoFromBinaryFile(param_file, param))
|
||||
<< "Failed to parse GraphDef file: " << param_file;
|
||||
}
|
||||
|
||||
void ReadTFNetParamsFromBinaryBufferOrDie(const char* data, size_t len,
|
||||
tensorflow::GraphDef* param) {
|
||||
CHECK(ReadProtoFromBinaryBuffer(data, len, param))
|
||||
<< "Failed to parse GraphDef buffer";
|
||||
}
|
||||
|
||||
void ReadTFNetParamsFromTextFileOrDie(const char* param_file,
|
||||
tensorflow::GraphDef* param) {
|
||||
CHECK(ReadProtoFromTextFileTF(param_file, param))
|
||||
<< "Failed to parse GraphDef file: " << param_file;
|
||||
CHECK(ReadProtoFromTextFile(param_file, param))
|
||||
<< "Failed to parse GraphDef file: " << param_file;
|
||||
}
|
||||
|
||||
void ReadTFNetParamsFromTextBufferOrDie(const char* data, size_t len,
|
||||
tensorflow::GraphDef* param) {
|
||||
CHECK(ReadProtoFromTextBuffer(data, len, param))
|
||||
<< "Failed to parse GraphDef buffer";
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -25,6 +25,13 @@ void ReadTFNetParamsFromBinaryFileOrDie(const char* param_file,
|
||||
void ReadTFNetParamsFromTextFileOrDie(const char* param_file,
|
||||
tensorflow::GraphDef* param);
|
||||
|
||||
// Read parameters from a memory buffer into a GraphDef proto message.
|
||||
void ReadTFNetParamsFromBinaryBufferOrDie(const char* data, size_t len,
|
||||
tensorflow::GraphDef* param);
|
||||
|
||||
void ReadTFNetParamsFromTextBufferOrDie(const char* data, size_t len,
|
||||
tensorflow::GraphDef* param);
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user