File size: 2,514 Bytes
2d9336d
 
 
 
 
 
0f53d40
2d9336d
 
 
 
 
 
0f53d40
2d9336d
0f53d40
 
 
 
2d9336d
 
 
 
 
 
 
0f53d40
 
 
 
 
 
 
 
 
 
 
 
 
 
8a919a4
 
2d9336d
 
 
 
 
 
0f53d40
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d9336d
 
 
 
 
 
 
0f53d40
 
 
 
 
 
 
 
 
 
 
 
59a7be2
8a919a4
2d9336d
 
 
 
8a919a4
2d9336d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8a919a4
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114

# app/core/llm_engine.py

import os

from langchain_nvidia_ai_endpoints import ChatNVIDIA
from langchain_groq import ChatGroq

# ============================================================
# Configuration
# ============================================================

NVIDIA_API_KEY = os.getenv("NVIDIA_API_KEY")
GROQ_API_KEY= os.getenv("GROQ_API_KEY")

MAIN_MODEL= "meta/llama-3.1-70b-instruct"
MAIN_MODEL_GROQ= "llama-3.3-70b-versatile"
EVAL_MODEL= "meta/llama-3.1-8b-instruct"
EVAL_MODEL_GROQ= "llama-3.1-8b-instant"


# ============================================================
# Main LLM (Non-streaming)
# Used everywhere graph.invoke() is still called.
# ============================================================

# llm = ChatNVIDIA(
#     model=MAIN_MODEL,
#     api_key=NVIDIA_API_KEY,
#     temperature=0.7,
#     max_tokens=1024,
# )


llm = ChatGroq(
            model=MAIN_MODEL_GROQ,
            api_key=GROQ_API_KEY,
            temperature=0.7,
            max_tokens=1024,
            # timeout=120,
)

# ============================================================
# Streaming LLM
# Used by /query-stream endpoint.
# Supports .astream()
# ============================================================

# streaming_llm = ChatNVIDIA(
#     model=MAIN_MODEL,
#     api_key=NVIDIA_API_KEY,
#     temperature=0.7,
#     max_tokens=1024,
#     streaming=True,
# )


streaming_llm= ChatGroq(
            model=MAIN_MODEL_GROQ,
            api_key=GROQ_API_KEY,
            temperature=0.7,
            max_tokens=1024,
            # timeout=120,
)

# ============================================================
# Evaluator LLM
# Faster + deterministic
# ============================================================

# eval_llm = ChatNVIDIA(
#     model=EVAL_MODEL,
#     api_key=NVIDIA_API_KEY,
#     temperature=0.0,
#     max_tokens=200,
# )

eval_llm= ChatGroq(
            model=EVAL_MODEL_GROQ,
            api_key=GROQ_API_KEY,
            temperature=0.0,
            max_tokens=200,
)

# ============================================================
# Helper getters
# (optional, but keeps imports clean)
# ============================================================

def get_llm():
    """
    Standard synchronous LLM.
    """
    return llm


def get_streaming_llm():
    """
    Streaming LLM.
    Use with:

        async for chunk in llm.astream(...):
            ...
    """
    return streaming_llm


def get_eval_llm():
    """
    Evaluator model.
    """
    return eval_llm