MCPcopy Create free account
hub / github.com/MrNeRF/LichtFeld-Studio / TEST_F

Function TEST_F

tests/test_tensor_compat.cpp:104–134  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

102// ============= Complex Expression Tests =============
103
104TEST_F(TensorTorchCompatTest, ComplexExpression1) {
105 // Test: (a + b) * c - d / e
106 std::vector<float> data_a(12), data_b(12), data_c(12), data_d(12), data_e(12);
107 for (auto& val : data_a)
108 val = dist(gen);
109 for (auto& val : data_b)
110 val = dist(gen);
111 for (auto& val : data_c)
112 val = dist(gen);
113 for (auto& val : data_d)
114 val = dist(gen);
115 for (auto& val : data_e)
116 val = std::abs(dist(gen)) + 1.0f; // Avoid division by zero
117
118 auto custom_a = Tensor::from_vector(data_a, {3, 4}, Device::CUDA);
119 auto custom_b = Tensor::from_vector(data_b, {3, 4}, Device::CUDA);
120 auto custom_c = Tensor::from_vector(data_c, {3, 4}, Device::CUDA);
121 auto custom_d = Tensor::from_vector(data_d, {3, 4}, Device::CUDA);
122 auto custom_e = Tensor::from_vector(data_e, {3, 4}, Device::CUDA);
123
124 auto torch_a = torch::tensor(data_a, torch::TensorOptions().device(torch::kCUDA)).reshape({3, 4});
125 auto torch_b = torch::tensor(data_b, torch::TensorOptions().device(torch::kCUDA)).reshape({3, 4});
126 auto torch_c = torch::tensor(data_c, torch::TensorOptions().device(torch::kCUDA)).reshape({3, 4});
127 auto torch_d = torch::tensor(data_d, torch::TensorOptions().device(torch::kCUDA)).reshape({3, 4});
128 auto torch_e = torch::tensor(data_e, torch::TensorOptions().device(torch::kCUDA)).reshape({3, 4});
129
130 auto custom_result = (custom_a + custom_b) * custom_c - custom_d / custom_e;
131 auto torch_result = (torch_a + torch_b) * torch_c - torch_d / torch_e;
132
133 compare_tensors(custom_result, torch_result, 1e-4f, 1e-5f, "ComplexExpression1");
134}
135
136TEST_F(TensorTorchCompatTest, ComplexExpression2) {
137 // Test: sigmoid(a * 2 + b) * relu(c - 1)

Callers

nothing calls this directly

Calls 15

from_vectorFunction · 0.85
TensorOptionsClass · 0.85
logFunction · 0.85
fullFunction · 0.85
sumFunction · 0.85
sigmoidMethod · 0.80
reluMethod · 0.80
expMethod · 0.80
viewMethod · 0.80
sum_scalarMethod · 0.80
mean_scalarMethod · 0.80
min_scalarMethod · 0.80

Tested by

no test coverage detected