MCPcopy Create free account
hub / github.com/OpenPipe/OpenPipe / LoraBaseModel

Class LoraBaseModel

trainer/src/lora_inference/main.py:85–161  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

83 container_idle_timeout=300,
84)
85class LoraBaseModel:
86 def __init__(self, base_model_id: str):
87 model_dir = merged_model_cache_dir(base_model_id)
88 logging.info(f"Loading base model {model_dir}")
89 self.seen_models = set()
90
91 self.engine = AsyncLLMEngine.from_engine_args(
92 AsyncEngineArgs(
93 model=base_model_id,
94 enable_lora=True,
95 max_loras=MAX_INPUTS,
96 max_lora_rank=8,
97 max_cpu_loras=500,
98 )
99 )
100
101 @modal.method()
102 async def generate(self, request: Input) -> Output:
103 logging.info(f"Processing for model {request.lora_model}")
104 lora_dir = lora_model_cache_dir(request.lora_model)
105
106 if request.lora_model not in self.seen_models:
107 if not os.path.exists(lora_model_cache_dir(request.lora_model)):
108 logging.info(f"Couldn't find model, reloading {lora_dir}")
109 stub.volume.reload()
110
111 if not os.path.exists(lora_model_cache_dir(request.lora_model)):
112 raise Exception(f"Couldn't find model {lora_dir}after reloading!")
113
114 self.seen_models.add(request.lora_model)
115
116 sample_params = SamplingParams(
117 n=request.n,
118 temperature=request.temperature,
119 max_tokens=request.max_tokens,
120 )
121
122 lora_request = LoRARequest(
123 request.lora_model,
124 abs(hash(request.lora_model)),
125 lora_dir,
126 )
127
128 request_id = random_uuid()
129
130 logging.info(f"Generating for request {request_id}")
131 output_generator = self.engine.generate(
132 request.prompt,
133 sample_params,
134 request_id=request_id,
135 lora_request=lora_request,
136 )
137
138 final_output: Union[RequestOutput, None] = None
139 async for request_output in output_generator:
140 # TODO: support streaming
141 final_output = request_output
142

Callers 1

chat_completionFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected