Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
test_llm.py278 linesDownload Raw Back to huggingface_hub
1import os2from collections.abc import Generator3 4import pytest5 6from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta7from core.model_runtime.entities.message_entities import AssistantPromptMessage, UserPromptMessage8from core.model_runtime.errors.validate import CredentialsValidateFailedError9from core.model_runtime.model_providers.huggingface_hub.llm.llm import HuggingfaceHubLargeLanguageModel10from tests.integration_tests.model_runtime.__mock.huggingface import setup_huggingface_mock11 12 13@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)14def test_hosted_inference_api_validate_credentials(setup_huggingface_mock):15    model = HuggingfaceHubLargeLanguageModel()16 17    with pytest.raises(CredentialsValidateFailedError):18        model.validate_credentials(19            model="HuggingFaceH4/zephyr-7b-beta",20            credentials={"huggingfacehub_api_type": "hosted_inference_api", "huggingfacehub_api_token": "invalid_key"},21        )22 23    with pytest.raises(CredentialsValidateFailedError):24        model.validate_credentials(25            model="fake-model",26            credentials={"huggingfacehub_api_type": "hosted_inference_api", "huggingfacehub_api_token": "invalid_key"},27        )28 29    model.validate_credentials(30        model="HuggingFaceH4/zephyr-7b-beta",31        credentials={32            "huggingfacehub_api_type": "hosted_inference_api",33            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),34        },35    )36 37 38@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)39def test_hosted_inference_api_invoke_model(setup_huggingface_mock):40    model = HuggingfaceHubLargeLanguageModel()41 42    response = model.invoke(43        model="HuggingFaceH4/zephyr-7b-beta",44        credentials={45            "huggingfacehub_api_type": "hosted_inference_api",46            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),47        },48        prompt_messages=[UserPromptMessage(content="Who are you?")],49        model_parameters={50            "temperature": 1.0,51            "top_k": 2,52            "top_p": 0.5,53        },54        stop=["How"],55        stream=False,56        user="abc-123",57    )58 59    assert isinstance(response, LLMResult)60    assert len(response.message.content) > 061 62 63@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)64def test_hosted_inference_api_invoke_stream_model(setup_huggingface_mock):65    model = HuggingfaceHubLargeLanguageModel()66 67    response = model.invoke(68        model="HuggingFaceH4/zephyr-7b-beta",69        credentials={70            "huggingfacehub_api_type": "hosted_inference_api",71            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),72        },73        prompt_messages=[UserPromptMessage(content="Who are you?")],74        model_parameters={75            "temperature": 1.0,76            "top_k": 2,77            "top_p": 0.5,78        },79        stop=["How"],80        stream=True,81        user="abc-123",82    )83 84    assert isinstance(response, Generator)85 86    for chunk in response:87        assert isinstance(chunk, LLMResultChunk)88        assert isinstance(chunk.delta, LLMResultChunkDelta)89        assert isinstance(chunk.delta.message, AssistantPromptMessage)90        assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True91 92 93@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)94def test_inference_endpoints_text_generation_validate_credentials(setup_huggingface_mock):95    model = HuggingfaceHubLargeLanguageModel()96 97    with pytest.raises(CredentialsValidateFailedError):98        model.validate_credentials(99            model="openchat/openchat_3.5",100            credentials={101                "huggingfacehub_api_type": "inference_endpoints",102                "huggingfacehub_api_token": "invalid_key",103                "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),104                "task_type": "text-generation",105            },106        )107 108    model.validate_credentials(109        model="openchat/openchat_3.5",110        credentials={111            "huggingfacehub_api_type": "inference_endpoints",112            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),113            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),114            "task_type": "text-generation",115        },116    )117 118 119@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)120def test_inference_endpoints_text_generation_invoke_model(setup_huggingface_mock):121    model = HuggingfaceHubLargeLanguageModel()122 123    response = model.invoke(124        model="openchat/openchat_3.5",125        credentials={126            "huggingfacehub_api_type": "inference_endpoints",127            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),128            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),129            "task_type": "text-generation",130        },131        prompt_messages=[UserPromptMessage(content="Who are you?")],132        model_parameters={133            "temperature": 1.0,134            "top_k": 2,135            "top_p": 0.5,136        },137        stop=["How"],138        stream=False,139        user="abc-123",140    )141 142    assert isinstance(response, LLMResult)143    assert len(response.message.content) > 0144 145 146@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)147def test_inference_endpoints_text_generation_invoke_stream_model(setup_huggingface_mock):148    model = HuggingfaceHubLargeLanguageModel()149 150    response = model.invoke(151        model="openchat/openchat_3.5",152        credentials={153            "huggingfacehub_api_type": "inference_endpoints",154            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),155            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),156            "task_type": "text-generation",157        },158        prompt_messages=[UserPromptMessage(content="Who are you?")],159        model_parameters={160            "temperature": 1.0,161            "top_k": 2,162            "top_p": 0.5,163        },164        stop=["How"],165        stream=True,166        user="abc-123",167    )168 169    assert isinstance(response, Generator)170 171    for chunk in response:172        assert isinstance(chunk, LLMResultChunk)173        assert isinstance(chunk.delta, LLMResultChunkDelta)174        assert isinstance(chunk.delta.message, AssistantPromptMessage)175        assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True176 177 178@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)179def test_inference_endpoints_text2text_generation_validate_credentials(setup_huggingface_mock):180    model = HuggingfaceHubLargeLanguageModel()181 182    with pytest.raises(CredentialsValidateFailedError):183        model.validate_credentials(184            model="google/mt5-base",185            credentials={186                "huggingfacehub_api_type": "inference_endpoints",187                "huggingfacehub_api_token": "invalid_key",188                "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),189                "task_type": "text2text-generation",190            },191        )192 193    model.validate_credentials(194        model="google/mt5-base",195        credentials={196            "huggingfacehub_api_type": "inference_endpoints",197            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),198            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),199            "task_type": "text2text-generation",200        },201    )202 203 204@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)205def test_inference_endpoints_text2text_generation_invoke_model(setup_huggingface_mock):206    model = HuggingfaceHubLargeLanguageModel()207 208    response = model.invoke(209        model="google/mt5-base",210        credentials={211            "huggingfacehub_api_type": "inference_endpoints",212            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),213            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),214            "task_type": "text2text-generation",215        },216        prompt_messages=[UserPromptMessage(content="Who are you?")],217        model_parameters={218            "temperature": 1.0,219            "top_k": 2,220            "top_p": 0.5,221        },222        stop=["How"],223        stream=False,224        user="abc-123",225    )226 227    assert isinstance(response, LLMResult)228    assert len(response.message.content) > 0229 230 231@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)232def test_inference_endpoints_text2text_generation_invoke_stream_model(setup_huggingface_mock):233    model = HuggingfaceHubLargeLanguageModel()234 235    response = model.invoke(236        model="google/mt5-base",237        credentials={238            "huggingfacehub_api_type": "inference_endpoints",239            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),240            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),241            "task_type": "text2text-generation",242        },243        prompt_messages=[UserPromptMessage(content="Who are you?")],244        model_parameters={245            "temperature": 1.0,246            "top_k": 2,247            "top_p": 0.5,248        },249        stop=["How"],250        stream=True,251        user="abc-123",252    )253 254    assert isinstance(response, Generator)255 256    for chunk in response:257        assert isinstance(chunk, LLMResultChunk)258        assert isinstance(chunk.delta, LLMResultChunkDelta)259        assert isinstance(chunk.delta.message, AssistantPromptMessage)260        assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True261 262 263def test_get_num_tokens():264    model = HuggingfaceHubLargeLanguageModel()265 266    num_tokens = model.get_num_tokens(267        model="google/mt5-base",268        credentials={269            "huggingfacehub_api_type": "inference_endpoints",270            "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),271            "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),272            "task_type": "text2text-generation",273        },274        prompt_messages=[UserPromptMessage(content="Hello World!")],275    )276 277    assert num_tokens == 7278