Skip to content

Commit 95050df

Browse files
fix: delete useless prompt_logprobs=1
1 parent 9e7f8a2 commit 95050df

2 files changed

Lines changed: 13 additions & 32 deletions

File tree

graphgen/bases/base_llm_wrapper.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,11 +26,11 @@ def __init__(
2626
**kwargs: Any,
2727
):
2828
self.system_prompt = system_prompt
29-
self.temperature = temperature
30-
self.max_tokens = max_tokens
31-
self.repetition_penalty = repetition_penalty
32-
self.top_p = top_p
33-
self.top_k = top_k
29+
self.temperature = float(temperature)
30+
self.max_tokens = int(max_tokens)
31+
self.repetition_penalty = float(repetition_penalty)
32+
self.top_p = float(top_p)
33+
self.top_k = int(top_k)
3434
self.tokenizer = tokenizer
3535

3636
for k, v in kwargs.items():

graphgen/models/llm/local/vllm_wrapper.py

Lines changed: 8 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -20,13 +20,9 @@ def __init__(
2020
temperature: float = 0.6,
2121
top_p: float = 1.0,
2222
top_k: int = 5,
23-
timeout: float = 300
23+
timeout: float = 300,
2424
**kwargs: Any,
2525
):
26-
temperature = float(temperature)
27-
top_p = float(top_p)
28-
top_k = int(top_k)
29-
3026
super().__init__(temperature=temperature, top_p=top_p, top_k=top_k, **kwargs)
3127
try:
3228
from vllm import AsyncEngineArgs, AsyncLLMEngine, SamplingParams
@@ -45,10 +41,7 @@ def __init__(
4541
disable_log_stats=False,
4642
)
4743
self.engine = AsyncLLMEngine.from_engine_args(engine_args)
48-
self.temperature = temperature
49-
self.top_p = top_p
50-
self.top_k = top_k
51-
self.timeout = timeout
44+
self.timeout = float(timeout)
5245

5346
@staticmethod
5447
def _build_inputs(prompt: str, history: Optional[List[str]] = None) -> str:
@@ -89,20 +82,15 @@ async def generate_answer(
8982
self._consume_generator(result_generator),
9083
timeout=self.timeout
9184
)
92-
85+
9386
if not final_output or not final_output.outputs:
9487
return ""
9588

9689
result_text = final_output.outputs[0].text
9790
return result_text
98-
99-
except asyncio.TimeoutError:
100-
await self.engine.abort(request_id)
101-
raise
102-
except asyncio.CancelledError:
103-
await self.engine.abort(request_id)
104-
raise
91+
10592
except Exception as e:
93+
print(f"Error in generate_answer: {e}")
10694
await self.engine.abort(request_id)
10795
raise
10896

@@ -116,7 +104,6 @@ async def generate_topk_per_token(
116104
temperature=0,
117105
max_tokens=1,
118106
logprobs=self.top_k,
119-
prompt_logprobs=1,
120107
)
121108

122109
result_generator = self.engine.generate(full_prompt, sp, request_id=request_id)
@@ -126,7 +113,7 @@ async def generate_topk_per_token(
126113
self._consume_generator(result_generator),
127114
timeout=self.timeout
128115
)
129-
116+
130117
if (
131118
not final_output
132119
or not final_output.outputs
@@ -154,14 +141,9 @@ async def generate_topk_per_token(
154141
)
155142
return [main_token]
156143
return []
157-
158-
except asyncio.TimeoutError:
159-
await self.engine.abort(request_id)
160-
raise
161-
except asyncio.CancelledError:
162-
await self.engine.abort(request_id)
163-
raise
144+
164145
except Exception as e:
146+
print(f"Error in generate_topk_per_token: {e}")
165147
await self.engine.abort(request_id)
166148
raise
167149

@@ -171,4 +153,3 @@ async def generate_inputs_prob(
171153
raise NotImplementedError(
172154
"VLLMWrapper does not support per-token logprobs yet."
173155
)
174-

0 commit comments

Comments
 (0)