mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-10-01 21:33:22 +08:00
59 lines
2.0 KiB
Python
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
|