Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 9 additions & 7 deletions lit_nlp/components/lime_explainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,11 +177,11 @@ def run(
config_defaults = {k: v.default for k, v in self.config_spec().items()}
config = dict(config_defaults, **(config or {})) # update and return

kernel_width = int(config[KERNEL_WIDTH_KEY])
num_samples = int(config[NUM_SAMPLES_KEY])
kernel_width = int(config[KERNEL_WIDTH_KEY]) # pyrefly: ignore[bad-argument-type]
num_samples = int(config[NUM_SAMPLES_KEY]) # pyrefly: ignore[bad-argument-type]
mask_string = (config[MASK_KEY])
# pylint: disable=g-explicit-bool-comparison
seed = int(config[SEED_KEY]) if config[SEED_KEY] != '' else None
seed = int(config[SEED_KEY]) if config[SEED_KEY] != '' else None # pyrefly: ignore[bad-argument-type]
# pylint: enable=g-explicit-bool-comparison

# Find keys of input (text) segments to explain.
Expand All @@ -201,7 +201,7 @@ def run(
return None

if (field := config[TARGET_HEAD_KEY]) and (
cls_idx := int(config[CLASS_KEY])
cls_idx := int(config[CLASS_KEY]) # pyrefly: ignore[bad-argument-type]
) != -1:
# TODO(b/205996131): remove this case
pred_key = field
Expand Down Expand Up @@ -249,10 +249,12 @@ def run(
class_to_explain=class_to_explain,
num_samples=num_samples,
tokenizer=str.split,
mask_token=mask_string,
mask_token=mask_string, # pyrefly: ignore[bad-argument-type]
kernel=functools.partial(
lime.exponential_kernel, kernel_width=kernel_width),
seed=seed)
lime.exponential_kernel, kernel_width=kernel_width
),
seed=seed,
)

# Turn the LIME explanation into a list following original word order.
scores = explanation.feature_importance
Expand Down
5 changes: 3 additions & 2 deletions lit_nlp/components/remote_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,9 +93,10 @@ def predict_minibatch(self, inputs: list[JsonDict]) -> list[JsonDict]:
'get_preds',
params={ # pyrefly: ignore[bad-argument-type]
'model': self._name,
'response_simple_json': False
'response_simple_json': False, # pyrefly: ignore[bad-assignment]
},
inputs=indexed_inputs)
inputs=indexed_inputs,
)
logging.info('Received %d predictions from remote model.', len(preds))
return preds

Expand Down
2 changes: 1 addition & 1 deletion lit_nlp/components/shap_explainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@ def run(

example_data = inputs or dataset.examples
examples: pd.DataFrame = pd.DataFrame(example_data)[input_feats]
sample_size = int(config.get(SAMPLE_KEY, 0))
sample_size = int(config.get(SAMPLE_KEY, 0)) # pyrefly: ignore[bad-argument-type]
if sample_size and len(examples) > sample_size:
inputs_to_use: pd.DataFrame = examples.sample(sample_size)
else:
Expand Down
4 changes: 2 additions & 2 deletions lit_nlp/examples/glue/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -513,7 +513,7 @@ def predict_minibatch(self, inputs: Iterable[JsonDict]):
)
# <float32>[batch_size, 1]
denom = tf.reduce_sum(token_mask, axis=1)
for i, layer_output in enumerate(out.hidden_states): # pyrefly: ignore[bad-argument-type]
for i, layer_output in enumerate(out.hidden_states): # pyrefly: ignore[bad-argument-type, not-iterable]
# layer_output is <float32>[batch_size, num_tokens, emb_dim]
# average over tokens to get <float32>[batch_size, emb_dim]
batched_outputs[f"layer_{i}/avg_emb"] = (
Expand All @@ -529,7 +529,7 @@ def predict_minibatch(self, inputs: Iterable[JsonDict]):
f"{self.model.config.num_hidden_layers}, got "
f"{len(out.attentions)}."
)
for i, layer_attention in enumerate(out.attentions): # pyrefly: ignore[bad-argument-type]
for i, layer_attention in enumerate(out.attentions): # pyrefly: ignore[bad-argument-type, not-iterable]
batched_outputs[f"layer_{i+1}/attention"] = layer_attention

if self.is_regression:
Expand Down
Loading