Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
088daf78d4 |
@@ -335,13 +335,13 @@ class Pipeline(_ScikitCompat):
|
||||
self.tokenizer = tokenizer
|
||||
self.modelcard = modelcard
|
||||
self.framework = framework
|
||||
self.device = device
|
||||
self.device = device if framework == "tf" else torch.device("cpu" if device < 0 else f"cuda:{device}")
|
||||
self.binary_output = binary_output
|
||||
self._args_parser = args_parser or DefaultArgumentHandler()
|
||||
|
||||
# Special handling
|
||||
if self.device >= 0 and self.framework == "pt":
|
||||
self.model = self.model.to("cuda:{}".format(self.device))
|
||||
if self.framework == "pt" and self.device.index >= 0:
|
||||
self.model = self.model.to(self.device)
|
||||
|
||||
def save_pretrained(self, save_directory):
|
||||
"""
|
||||
@@ -385,11 +385,19 @@ class Pipeline(_ScikitCompat):
|
||||
with tf.device("/CPU:0" if self.device == -1 else "/device:GPU:{}".format(self.device)):
|
||||
yield
|
||||
else:
|
||||
if self.device >= 0:
|
||||
if self.device.index >= 0:
|
||||
torch.cuda.set_device(self.device)
|
||||
|
||||
yield
|
||||
|
||||
def ensure_tensor_on_device(self, **inputs):
|
||||
"""
|
||||
Ensure PyTorch tensors are on the specified device.
|
||||
:param inputs:
|
||||
:return:
|
||||
"""
|
||||
return {name: tensor.to(self.device) for name, tensor in inputs.items()}
|
||||
|
||||
def inputs_for_model(self, features: Union[dict, List[dict]]) -> Dict:
|
||||
"""
|
||||
Generates the input dictionary with model-specific parameters.
|
||||
@@ -415,16 +423,13 @@ class Pipeline(_ScikitCompat):
|
||||
def __call__(self, *texts, **kwargs):
|
||||
# Parse arguments
|
||||
inputs = self._args_parser(*texts, **kwargs)
|
||||
inputs = self.tokenizer.batch_encode_plus(
|
||||
inputs, add_special_tokens=True, return_tensors=self.framework, max_length=self.tokenizer.max_len
|
||||
)
|
||||
|
||||
# Encode for forward
|
||||
with self.device_placement():
|
||||
inputs = self.tokenizer.batch_encode_plus(
|
||||
inputs, add_special_tokens=True, return_tensors=self.framework, max_length=self.tokenizer.max_len
|
||||
)
|
||||
|
||||
# Filter out features not available on specific models
|
||||
inputs = self.inputs_for_model(inputs)
|
||||
return self._forward(inputs)
|
||||
# Filter out features not available on specific models
|
||||
inputs = self.inputs_for_model(inputs)
|
||||
return self._forward(inputs)
|
||||
|
||||
def _forward(self, inputs):
|
||||
"""
|
||||
@@ -434,12 +439,15 @@ class Pipeline(_ScikitCompat):
|
||||
Returns:
|
||||
Numpy array
|
||||
"""
|
||||
if self.framework == "tf":
|
||||
# TODO trace model
|
||||
predictions = self.model(inputs, training=False)[0]
|
||||
else:
|
||||
with torch.no_grad():
|
||||
predictions = self.model(**inputs)[0].cpu()
|
||||
# Encode for forward
|
||||
with self.device_placement():
|
||||
if self.framework == "tf":
|
||||
# TODO trace model
|
||||
predictions = self.model(inputs, training=False)[0]
|
||||
else:
|
||||
with torch.no_grad():
|
||||
inputs = self.ensure_tensor_on_device(**inputs)
|
||||
predictions = self.model(**inputs)[0].cpu()
|
||||
|
||||
return predictions.numpy()
|
||||
|
||||
@@ -534,6 +542,7 @@ class NerPipeline(Pipeline):
|
||||
input_ids = tokens["input_ids"].numpy()[0]
|
||||
else:
|
||||
with torch.no_grad():
|
||||
tokens = self.ensure_tensor_on_device(**tokens)
|
||||
entities = self.model(**tokens)[0][0].cpu().numpy()
|
||||
input_ids = tokens["input_ids"].cpu().numpy()[0]
|
||||
|
||||
@@ -710,7 +719,7 @@ class QuestionAnsweringPipeline(Pipeline):
|
||||
else:
|
||||
with torch.no_grad():
|
||||
# Retrieve the score for the context tokens only (removing question tokens)
|
||||
fw_args = {k: torch.tensor(v) for (k, v) in fw_args.items()}
|
||||
fw_args = {k: torch.tensor(v, device=self.device) for (k, v) in fw_args.items()}
|
||||
start, end = self.model(**fw_args)
|
||||
start, end = start.cpu().numpy(), end.cpu().numpy()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user