Merge pull request #274 from makaveli10/fallback_to_fp32
Set compute_type based on device capability.
This commit is contained in:
@@ -787,9 +787,15 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
self.no_speech_thresh = 0.45
|
self.no_speech_thresh = 0.45
|
||||||
|
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
if device == "cuda":
|
||||||
|
major, _ = torch.cuda.get_device_capability(device)
|
||||||
|
self.compute_type = "float16" if major >= 7 else "float32"
|
||||||
|
else:
|
||||||
|
self.compute_type = "int8"
|
||||||
|
|
||||||
if self.model_size_or_path is None:
|
if self.model_size_or_path is None:
|
||||||
return
|
return
|
||||||
|
logging.info(f"Using Device={device} with precision {self.compute_type}")
|
||||||
|
|
||||||
if single_model:
|
if single_model:
|
||||||
if ServeClientFasterWhisper.SINGLE_MODEL is None:
|
if ServeClientFasterWhisper.SINGLE_MODEL is None:
|
||||||
@@ -822,7 +828,7 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
self.transcriber = WhisperModel(
|
self.transcriber = WhisperModel(
|
||||||
self.model_size_or_path,
|
self.model_size_or_path,
|
||||||
device=device,
|
device=device,
|
||||||
compute_type="int8" if device == "cpu" else "float16",
|
compute_type=self.compute_type,
|
||||||
local_files_only=False,
|
local_files_only=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user