项目文件夹

文件
wehub-resource-sync 593b94c120
pytest / Unit Tests (push) Has been cancelled
pytest / Integration (integration_tests_a) (push) Has been cancelled
pytest / Integration (integration_tests_b) (push) Has been cancelled
pytest / Integration (integration_tests_c) (push) Has been cancelled
pytest / Integration (integration_tests_d) (push) Has been cancelled
pytest / Integration (integration_tests_e) (push) Has been cancelled
pytest / Integration (integration_tests_f) (push) Has been cancelled
pytest / Integration (integration_tests_g) (push) Has been cancelled
pytest / Integration (integration_tests_h) (push) Has been cancelled
pytest / Integration (integration_tests_i) (push) Has been cancelled
pytest / Integration (integration_tests_j) (push) Has been cancelled
pytest / Distributed (distributed_a) (push) Has been cancelled
pytest / Distributed (distributed_b) (push) Has been cancelled
pytest / Distributed (distributed_c) (push) Has been cancelled
pytest / Distributed (distributed_d) (push) Has been cancelled
pytest / Distributed (distributed_e) (push) Has been cancelled
pytest / Distributed (distributed_f) (push) Has been cancelled
pytest / Minimal Install (push) Has been cancelled
pytest / Event File (push) Has been cancelled
pytest (slow) / py-slow (push) Has been cancelled
Publish JSON Schema / publish-schema (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:49:20 +08:00

92 行
3.2 KiB
Python

import pytest
import torch
from ludwig.constants import ENCODER, ENCODER_OUTPUT
from ludwig.features.binary_feature import BinaryInputFeature, BinaryOutputFeature
from ludwig.schema.features.binary_feature import BinaryInputFeatureConfig, BinaryOutputFeatureConfig
from ludwig.schema.utils import load_config_with_kwargs
from ludwig.utils.torch_utils import get_torch_device
BATCH_SIZE = 2
BINARY_W_SIZE = 1
DEVICE = get_torch_device()
@pytest.fixture(scope="module")
def binary_config():
return {
"name": "binary_feature",
"type": "binary",
}
@pytest.mark.parametrize("encoder", ["passthrough", "dense"])
def test_binary_input_feature(binary_config: dict, encoder: str):
binary_config.update({ENCODER: {"type": encoder}})
binary_config, _ = load_config_with_kwargs(BinaryInputFeatureConfig, binary_config)
binary_input_feature = BinaryInputFeature(binary_config).to(DEVICE)
binary_tensor = binary_input_feature.create_sample_input(batch_size=BATCH_SIZE).to(DEVICE)
assert binary_tensor.shape == torch.Size([BATCH_SIZE])
assert binary_tensor.dtype == torch.bool
encoder_output = binary_input_feature(binary_tensor)
assert encoder_output[ENCODER_OUTPUT].shape[1:] == binary_input_feature.output_shape
def test_binary_output_feature():
binary_output_config = {
"name": "binary_feature",
"type": "binary",
"input_size": BINARY_W_SIZE,
"decoder": {
"type": "regressor",
"input_size": 1,
},
"loss": {
"type": "binary_weighted_cross_entropy",
"positive_class_weight": 1,
"robust_lambda": 0,
"confidence_penalty": 0,
},
}
binary_output_config, _ = load_config_with_kwargs(BinaryOutputFeatureConfig, binary_output_config)
binary_output_feature = BinaryOutputFeature(binary_output_config, {}).to(DEVICE)
combiner_outputs = dict()
combiner_outputs["combiner_output"] = torch.randn([BATCH_SIZE, BINARY_W_SIZE], dtype=torch.float32).to(DEVICE)
binary_output = binary_output_feature(combiner_outputs, {})
assert "last_hidden" in binary_output
assert "logits" in binary_output
assert binary_output["logits"].size() == torch.Size([BATCH_SIZE])
def test_binary_output_feature_without_positive_class_weight():
binary_output_config = {
"name": "binary_feature",
"type": "binary",
"input_size": BINARY_W_SIZE,
"decoder": {
"type": "regressor",
"input_size": 1,
},
"loss": {
"type": "binary_weighted_cross_entropy",
"positive_class_weight": None,
"robust_lambda": 0,
"confidence_penalty": 0,
},
}
binary_output_config, _ = load_config_with_kwargs(BinaryOutputFeatureConfig, binary_output_config)
binary_output_feature = BinaryOutputFeature(binary_output_config, {}).to(DEVICE)
combiner_outputs = {}
combiner_outputs["combiner_output"] = torch.randn([BATCH_SIZE, BINARY_W_SIZE], dtype=torch.float32).to(DEVICE)
binary_output = binary_output_feature(combiner_outputs, {})
assert "last_hidden" in binary_output
assert "logits" in binary_output
assert binary_output["logits"].size() == torch.Size([BATCH_SIZE])