"""LLM 显存保护逻辑的轻量测试(不需要下载模型)。""" from types import MethodType, SimpleNamespace from unittest.mock import patch import torch from lnp_ml.modeling.layers.llm_prompt import LLMPromptEncoder, _positive_int_env def _empty_encoder(micro_batch: int = 2) -> LLMPromptEncoder: encoder = LLMPromptEncoder.__new__(LLMPromptEncoder) torch.nn.Module.__init__(encoder) encoder.inference_batch_size = micro_batch return encoder def test_invalid_inference_batch_env_falls_back() -> None: with patch.dict("os.environ", {"LLM_INFERENCE_BATCH_SIZE": "invalid"}): assert _positive_int_env("LLM_INFERENCE_BATCH_SIZE", 4) == 4 with patch.dict("os.environ", {"LLM_INFERENCE_BATCH_SIZE": "0"}): assert _positive_int_env("LLM_INFERENCE_BATCH_SIZE", 4) == 4 def test_softrag_inference_is_microbatched() -> None: encoder = _empty_encoder(micro_batch=2).eval() calls = [] def fake_chunk(self, smiles, chem, tab, device): calls.append((list(smiles), chem.clone(), tab.clone())) return chem[:, 0, :] + tab[:, 0, :] encoder._encode_softrag_chunk = MethodType(fake_chunk, encoder) chem = torch.arange(5 * 2 * 3, dtype=torch.float32).reshape(5, 2, 3) tab = torch.ones(5, 1, 3) with torch.no_grad(): output = encoder._encode_softrag( [f"SMILES-{i}" for i in range(5)], chem, tab, torch.device("cpu") ) assert [len(smiles) for smiles, _, _ in calls] == [2, 2, 1] assert torch.equal(output, chem[:, 0, :] + tab[:, 0, :]) def test_qwen_forward_disables_kv_cache() -> None: encoder = _empty_encoder() encoder._is_qwen = True received = {} class FakeBackbone(torch.nn.Module): def forward(self, **kwargs): received.update(kwargs) return SimpleNamespace(last_hidden_state=torch.zeros(1, 1, 1)) encoder.encoder = FakeBackbone() encoder._encoder_forward(input_ids=torch.ones(1, 1, dtype=torch.long)) assert received["use_cache"] is False