lnp_ml/tests/test_llm_memory.py

59 lines
2.0 KiB
Python

"""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