MCPcopy Create free account
hub / github.com/Tencent/TurboTransformers / CheckCppBert

Function CheckCppBert

example/cpp/bert_model_test.cpp:27–58  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

25namespace turbo_transformers {
26namespace loaders {
27bool CheckCppBert(bool use_cuda, bool only_input) {
28 BertModel model(model_file_path,
29 use_cuda ? DLDeviceType::kDLGPU : DLDeviceType::kDLCPU, 12,
30 12);
31 std::vector<std::vector<int64_t>> position_ids{{1, 0, 0, 0}, {1, 1, 1, 0}};
32 std::vector<std::vector<int64_t>> segment_ids{{1, 1, 1, 0}, {1, 0, 0, 0}};
33 if (only_input) {
34 position_ids.clear();
35 segment_ids.clear();
36 }
37 auto vec = model({{12166, 10699, 16752, 4454}, {5342, 16471, 817, 16022}},
38 position_ids, segment_ids, PoolType::kFirst, false);
39 REQUIRE(vec.size() == 768 * 2);
40 // Write a better UT
41 for (size_t i = 0; i < vec.size(); ++i) {
42 REQUIRE(!std::isnan(vec.data()[i]));
43 REQUIRE(!std::isinf(vec.data()[i]));
44 }
45 if (only_input) {
46 std::cerr << vec.data()[0] << std::endl;
47 REQUIRE(fabs(vec.data()[0] - -0.9791) < 1e-3);
48 REQUIRE(fabs(vec.data()[1] - 0.8283) < 1e-3);
49 REQUIRE(fabs(vec.data()[768] - 0.3837) < 1e-3);
50 REQUIRE(fabs(vec.data()[768 + 1] - 0.1659) < 1e-3);
51 } else {
52 REQUIRE(fabs(vec.data()[0] - -0.3760) < 1e-3);
53 REQUIRE(fabs(vec.data()[1] - 0.5674) < 1e-3);
54 REQUIRE(fabs(vec.data()[768] - -0.2701) < 1e-3);
55 REQUIRE(fabs(vec.data()[768 + 1] - 0.2676) < 1e-3);
56 }
57 return true;
58}
59
60bool CheckCppBertWithPooler(bool use_cuda, bool only_input) {
61 BertModel model("models/bert.npz",

Callers 1

Calls 2

dataMethod · 0.80
clearMethod · 0.45

Tested by

no test coverage detected