@@ -159,9 +159,51 @@ class ReplicateImageRunner(ReplicateBaseRunner): return chunk -class ReplicateVideoRunner(BaseRunner): +class ReplicateVideoRunner(ReplicateImageRunner): @classmethod - def generate(cls, content=None, file=None, parameters={}, history=[], scrape_results=[]): ... + def generate(cls, content=None, file=None, parameters={}, history=[], scrape_results=[]): + model_key = parameters.pop('model_key', 'model') + model = parameters.pop(model_key) + model_owner = parameters.pop('model_owner') + official = model and model_owner + + payload = {} + if not official: + payload.update({'version': model}) + + payload['input'] = cls.compile_input( + content=content, + file=file, + parameters=parameters, + history=history, + scrape_results=scrape_results, + ) + + with httpx.Client( + base_url='https://api.replicate.com/v1', + headers={ + 'Authorization': f'Bearer {settings.REPLICATE_API_KEY}', + 'Content-Type': 'application/json', + }, + timeout=None, + ) as client: + resp = client.post( + f'/models/{model_owner}/{model}/predictions' if official else '/predictions', + json=payload + ) + data = resp.json() + while True: + resp = client.get(f"/predictions/{data['id']}") + result = resp.json() + logger.info(result) + if result['status'] in ('succeeded', 'failed'): + break + if result['status'] == 'failed': + raise Exception('Model failed generation') + stream_url = result['output'] + with client.stream('GET', stream_url) as stream: + for chunk in stream.iter_bytes(512 * 1024): + yield base64.b64encode(chunk).decode('utf-8') class ReplicateAudioRunner(BaseRunner):