Skip to content

Commit bcae735

Browse files
h-jooLIT team
authored andcommitted
Automated Code Change
PiperOrigin-RevId: 944116839
1 parent 871d9a5 commit bcae735

34 files changed

Lines changed: 143 additions & 141 deletions

lit_nlp/api/components.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -137,7 +137,7 @@ def is_compatible(self, model: lit_model.Model,
137137
"""True if the model and dataset support metric computation."""
138138
for pred_spec in model.output_spec().values():
139139
parent_key: Optional[str] = getattr(pred_spec, 'parent', None)
140-
parent_spec: Optional[types.LitType] = dataset.spec().get(parent_key)
140+
parent_spec: Optional[types.LitType] = dataset.spec().get(parent_key) # pyrefly: ignore[bad-argument-type]
141141
if self.is_field_compatible(pred_spec, parent_spec):
142142
return True
143143
return False

lit_nlp/api/dataset.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -236,7 +236,7 @@ def remap(self, field_map: Mapping[str, str]):
236236
"""Return a copy of this dataset with some fields renamed."""
237237
new_spec = utils.remap_dict(self.spec(), field_map)
238238
new_examples = [utils.remap_dict(ex, field_map) for ex in self.examples]
239-
return Dataset(new_spec, new_examples, base=self)
239+
return Dataset(new_spec, new_examples, base=self) # pyrefly: ignore[bad-argument-type]
240240

241241
@staticmethod
242242
def lit_example_from_bytes(input_bytes: bytes) -> Optional[JsonDict]:
@@ -273,7 +273,7 @@ def index_inputs(
273273
indexed.append(
274274
IndexedInput(
275275
data=types.MappingProxyType(
276-
example | {INPUT_ID_FIELD: ex_id, INPUT_META_FIELD: ex_meta}
276+
example | {INPUT_ID_FIELD: ex_id, INPUT_META_FIELD: ex_meta} # pyrefly: ignore[unsupported-operation]
277277
),
278278
id=ex_id,
279279
meta=ex_meta,
@@ -376,6 +376,7 @@ def load(self, path: str):
376376
new_dataset = base_dataset.load(path) if base_dataset else None
377377

378378
if new_dataset is not None:
379+
# pyrefly: ignore[missing-attribute]
379380
description = (f'{len(new_dataset)} examples from '
380381
f'{path}\n{self._base.description()}')
381382
return IndexedDataset(
@@ -438,7 +439,7 @@ def load_lit_format(
438439
**kw,
439440
)
440441
else:
441-
return Dataset(spec=spec, examples=examples, *args, **kw)
442+
return Dataset(spec=spec, examples=examples, *args, **kw) # pyrefly: ignore[bad-argument-type]
442443

443444

444445
# TODO(b/202210900): Remove "NoneDataset" once the LIT front-end constructs its

lit_nlp/api/layout.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -70,7 +70,7 @@ def __call__(self, **kw):
7070
# LINT.IfChange
7171
# TODO(lit-dev): consider making modules subclass this instead of LitModuleName.
7272
@attr.s(auto_attribs=True)
73-
class ModuleConfig(dtypes.DataTuple):
73+
class ModuleConfig(dtypes.DataTuple): # pyrefly: ignore[invalid-inheritance]
7474
module: Union[str, LitModuleName]
7575
requiredForTab: bool = False
7676
# TODO(b/172979677): support title, duplicateAsRow, numCols,
@@ -88,15 +88,15 @@ class ModuleConfig(dtypes.DataTuple):
8888

8989

9090
@attr.s(auto_attribs=True)
91-
class LayoutSettings(dtypes.DataTuple):
91+
class LayoutSettings(dtypes.DataTuple): # pyrefly: ignore[invalid-inheritance]
9292
hideToolbar: bool = False
9393
mainHeight: int = 45
9494
leftWidth: int = 50
9595
centerPage: bool = False
9696

9797

9898
@attr.s(auto_attribs=True)
99-
class LitCanonicalLayout(dtypes.DataTuple):
99+
class LitCanonicalLayout(dtypes.DataTuple): # pyrefly: ignore[invalid-inheritance]
100100
"""Frontend UI layout; should match client/lib/types.ts."""
101101
upper: LitTabGroupLayout
102102
lower: LitTabGroupLayout = attr.ib(factory=dict)

lit_nlp/api/model.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -178,12 +178,12 @@ def load(self, path: str):
178178
@abc.abstractmethod
179179
def input_spec(self) -> types.Spec:
180180
"""Return a spec describing model inputs."""
181-
return
181+
return # pyrefly: ignore[bad-return]
182182

183183
@abc.abstractmethod
184184
def output_spec(self) -> types.Spec:
185185
"""Return a spec describing model outputs."""
186-
return
186+
return # pyrefly: ignore[bad-return]
187187

188188
def get_embedding_table(self) -> tuple[list[str], np.ndarray]:
189189
"""Return the full vocabulary and embedding table.
@@ -369,7 +369,7 @@ def predict_minibatch(self, inputs: list[JsonDict]) -> list[JsonDict]:
369369
Returns:
370370
list of outputs, following model.output_spec()
371371
"""
372-
return
372+
return # pyrefly: ignore[bad-return]
373373

374374

375375
class ProjectorModel(BatchedModel, metaclass=abc.ABCMeta):

lit_nlp/api/types.py

Lines changed: 26 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -155,10 +155,10 @@ def from_json(d: JsonDict):
155155

156156
base_cls = globals().get("LitType")
157157
cls = globals().get(type_name) # class by name from this module
158-
if cls is None or not issubclass(cls, base_cls):
158+
if cls is None or not issubclass(cls, base_cls): # pyrefly: ignore[bad-argument-type]
159159
raise NameError(f"{type_name} is not a valid LitType.")
160160

161-
return cls(**{k: d[k] for k in d if k != "__name__"})
161+
return cls(**{k: d[k] for k in d if k != "__name__"}) # pyrefly: ignore[not-callable]
162162

163163
Spec = dict[str, LitType]
164164

@@ -174,7 +174,7 @@ def _remap_leaf(leaf: LitType, keymap: dict[str, str]) -> LitType:
174174
k: (keymap.get(v, v) if k in FIELD_REF_ATTRIBUTES else v)
175175
for k, v in d.items()
176176
}
177-
return leaf.__class__(**d)
177+
return leaf.__class__(**d) # pyrefly: ignore[bad-argument-type]
178178

179179

180180
def remap_spec(spec: Spec, keymap: dict[str, str]) -> Spec:
@@ -201,7 +201,7 @@ class StringLitType(LitType):
201201
Mainly used for string inputs that have special formatting, and should only
202202
be edited manually.
203203
"""
204-
default: str = ""
204+
default: str = "" # pyrefly: ignore[bad-override]
205205

206206
def validate_input(self, value, spec: Spec, example: Input):
207207
if not isinstance(value, str):
@@ -260,13 +260,13 @@ def validate_output(self, value, output_spec: Spec, output_dict: JsonDict,
260260
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
261261
class ListLitType(LitType):
262262
"""List type."""
263-
default: Sequence[Any] = None
263+
default: Sequence[Any] = None # pyrefly: ignore[bad-assignment, bad-override]
264264

265265

266266
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
267267
class _StringCandidateList(ListLitType):
268268
"""A list of (text, score) tuples."""
269-
default: ScoredTextCandidates = None
269+
default: ScoredTextCandidates = None # pyrefly: ignore[bad-assignment]
270270

271271
def validate_output(self, value, output_spec: Spec, output_dict: JsonDict,
272272
input_spec: Spec, dataset_spec: Spec,
@@ -373,9 +373,9 @@ class TokenTopKPreds(ListLitType):
373373
374374
The inner list should contain (word, probability) in descending order.
375375
"""
376-
default: Sequence[ScoredTextCandidates] = None
376+
default: Sequence[ScoredTextCandidates] = None # pyrefly: ignore[bad-assignment]
377377

378-
align: str = None # name of a Tokens field in the model output
378+
align: str = None # name of a Tokens field in the model output # pyrefly: ignore[bad-assignment]
379379
parent: Optional[str] = None
380380

381381
def _validate_scored_candidates(self, scored_candidates):
@@ -389,7 +389,7 @@ def _validate_scored_candidates(self, scored_candidates):
389389
if scored_candidate[1] is not None:
390390
if not isinstance(scored_candidate[1], NumericTypes):
391391
raise ValueError(f"{scored_candidate} second element is not a num")
392-
if prev_val < scored_candidate[1]:
392+
if prev_val < scored_candidate[1]: # pyrefly: ignore[unsupported-operation]
393393
raise ValueError(
394394
"TokenTopKPreds candidates are not in descending order")
395395
else:
@@ -414,7 +414,7 @@ class Scalar(LitType):
414414
"""Scalar value, a single float or int."""
415415
min_val: float = 0
416416
max_val: float = 1
417-
default: float = 0
417+
default: float = 0 # pyrefly: ignore[bad-override]
418418
step: float = .01
419419

420420
def validate_input(self, value, spec: Spec, example: Input):
@@ -464,7 +464,7 @@ class TokenScores(_FloatList):
464464
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
465465
class ReferenceScores(ListLitType):
466466
"""Score of one or more target sequences."""
467-
default: Sequence[float] = None
467+
default: Sequence[float] = None # pyrefly: ignore[bad-assignment]
468468

469469
# name of a TextSegment or ReferenceTexts field in the input
470470
parent: Optional[str] = None
@@ -504,7 +504,7 @@ def validate_input(self, value, spec: Spec, example: Input):
504504
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
505505
class _Tensor(LitType):
506506
"""A tensor type."""
507-
default: Sequence[float] = None
507+
default: Sequence[float] = None # pyrefly: ignore[bad-assignment, bad-override]
508508

509509
def validate_input(self, value, spec: Spec, example: Input):
510510
if isinstance(value, list):
@@ -585,7 +585,7 @@ class SpanLabels(ListLitType):
585585
Span labels can cover more than one token, may not cover all tokens in the
586586
sentence, and may overlap with each other.
587587
"""
588-
default: Sequence[dtypes.SpanLabel] = None
588+
default: Sequence[dtypes.SpanLabel] = None # pyrefly: ignore[bad-assignment]
589589
align: str # name of Tokens field
590590
parent: Optional[str] = None
591591

@@ -609,7 +609,7 @@ class EdgeLabels(ListLitType):
609609
https://github.com/nyu-mll/jiant/tree/master/probing#data-format for more
610610
details.
611611
"""
612-
default: Sequence[dtypes.EdgeLabel] = None
612+
default: Sequence[dtypes.EdgeLabel] = None # pyrefly: ignore[bad-assignment]
613613
align: str # name of Tokens field
614614

615615
def validate_output(self, value, output_spec: Spec, output_dict: JsonDict,
@@ -636,7 +636,7 @@ class MultiSegmentAnnotations(ListLitType):
636636
TODO(lit-dev): by default, spans are treated as bytes in this context.
637637
Make this configurable, if some spans need to refer to tokens instead.
638638
"""
639-
default: Sequence[dtypes.AnnotationCluster] = None
639+
default: Sequence[dtypes.AnnotationCluster] = None # pyrefly: ignore[bad-assignment]
640640
exclusive: bool = False # if true, treat as candidate list
641641
background: bool = False # if true, don't emphasize in visualization
642642

@@ -775,7 +775,7 @@ class SubwordOffsets(ListLitType):
775775
776776
offsets[i] should be the index of the first wordpiece for input token i.
777777
"""
778-
default: Sequence[int] = None
778+
default: Sequence[int] = None # pyrefly: ignore[bad-assignment]
779779
align_in: str # name of field in data spec
780780
align_out: str # name of field in model output spec
781781

@@ -793,7 +793,7 @@ class SparseMultilabelPreds(_StringCandidateList):
793793
794794
The tuples are of the label and the score.
795795
"""
796-
default: ScoredTextCandidates = None
796+
default: ScoredTextCandidates = None # pyrefly: ignore[bad-assignment]
797797
vocab: Optional[Sequence[str]] = None # label names
798798
parent: Optional[str] = None
799799

@@ -816,7 +816,7 @@ class SingleFieldMatcher(FieldMatcher):
816816
817817
UI will materialize this to a dropdown-list.
818818
"""
819-
default: str = None
819+
default: str = None # pyrefly: ignore[bad-assignment, bad-override]
820820

821821

822822
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
@@ -826,7 +826,7 @@ class MultiFieldMatcher(FieldMatcher):
826826
UI will materialize this to multiple checkboxes. Use this when the user needs
827827
to pick more than one field in UI.
828828
"""
829-
default: Sequence[str] = [] # default names of selected items.
829+
default: Sequence[str] = [] # default names of selected items. # pyrefly: ignore[bad-override]
830830
select_all: bool = False # Select all by default (overriddes default).
831831

832832

@@ -840,20 +840,20 @@ class Salience(LitType):
840840
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
841841
class TokenSalience(Salience):
842842
"""Metadata about a returned token salience map."""
843-
default: dtypes.TokenSalience = None
843+
default: dtypes.TokenSalience = None # pyrefly: ignore[bad-assignment, bad-override]
844844

845845

846846
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
847847
class FeatureSalience(Salience):
848848
"""Metadata about a returned feature salience map."""
849-
default: dtypes.FeatureSalience = None
849+
default: dtypes.FeatureSalience = None # pyrefly: ignore[bad-assignment, bad-override]
850850

851851

852852
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
853853
class FrameSalience(Salience):
854854
"""Metadata about a returned frame salience map."""
855855

856-
default: dtypes.FrameSalience = None
856+
default: dtypes.FrameSalience = None # pyrefly: ignore[bad-assignment, bad-override]
857857

858858

859859
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
@@ -869,13 +869,13 @@ class ImageSalience(Salience):
869869
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
870870
class SequenceSalience(Salience):
871871
"""Metadata about a returned sequence salience map."""
872-
default: dtypes.SequenceSalienceMap = None
872+
default: dtypes.SequenceSalienceMap = None # pyrefly: ignore[bad-assignment, bad-override]
873873

874874

875875
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
876876
class BooleanLitType(LitType):
877877
"""Boolean value."""
878-
default: bool = False
878+
default: bool = False # pyrefly: ignore[bad-override]
879879

880880
def validate_input(self, value, spec, example: Input):
881881
if not isinstance(value, bool):
@@ -916,7 +916,7 @@ class MetricBestValue(dtypes.EnumSerializableAsValues, enum.Enum):
916916
@attr.s(auto_attribs=True, frozen=True, kw_only=True)
917917
class MetricResult(LitType):
918918
"""Score returned from the computation of a Metric."""
919-
default: float = 0
919+
default: float = 0 # pyrefly: ignore[bad-override]
920920
description: str = ""
921921
best_value: MetricBestValue = MetricBestValue.NONE
922922

@@ -934,7 +934,7 @@ class SalienceTargetInfo(LitType):
934934
935935
Value is a dict with keys 'field' (str) and 'index' (Optional[int]).
936936
"""
937-
default: Optional[Mapping[str, Any]] = None
937+
default: Optional[Mapping[str, Any]] = None # pyrefly: ignore[bad-override]
938938

939939

940940
# LINT.ThenChange(../client/lib/lit_types.ts)

lit_nlp/app.py

Lines changed: 13 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -111,7 +111,7 @@ def _build_metadata(self):
111111
}
112112

113113
# List compatible datasets.
114-
info['datasets'] = [
114+
info['datasets'] = [ # pyrefly: ignore[bad-assignment]
115115
name for name, dataset in self._datasets.items()
116116
if model.is_compatible_with_dataset(dataset)
117117
]
@@ -134,13 +134,13 @@ def _build_metadata(self):
134134
_get_compatible_names(self._metrics, model, dataset)
135135
)
136136

137-
info['generators'] = [
137+
info['generators'] = [ # pyrefly: ignore[bad-assignment]
138138
name for name in self._generators.keys() if name in compat_gens
139139
]
140-
info['interpreters'] = [
140+
info['interpreters'] = [ # pyrefly: ignore[bad-assignment]
141141
name for name in self._interpreters.keys() if name in compat_interps
142142
]
143-
info['metrics'] = [
143+
info['metrics'] = [ # pyrefly: ignore[bad-assignment]
144144
name for name in self._metrics.keys() if name in compat_metrics
145145
]
146146
model_info[name] = info
@@ -213,7 +213,7 @@ def _reconstitute_inputs(
213213
len(inputs),
214214
dataset_name,
215215
)
216-
return [index[ex] if isinstance(ex, str) else ex for ex in inputs]
216+
return [index[ex] if isinstance(ex, str) else ex for ex in inputs] # pyrefly: ignore[bad-index]
217217

218218
def _save_datapoints(
219219
self,
@@ -297,16 +297,16 @@ def _get_preds(self,
297297

298298
# Figure out what to return to the frontend.
299299
output_spec = self._get_model_spec(model)['output']
300-
requested_types = requested_types.split(',') if requested_types else []
301-
requested_fields = requested_fields.split(',') if requested_fields else []
300+
requested_types = requested_types.split(',') if requested_types else [] # pyrefly: ignore[bad-assignment]
301+
requested_fields = requested_fields.split(',') if requested_fields else [] # pyrefly: ignore[bad-assignment]
302302
logging.info('Requested types: %s, fields: %s', str(requested_types),
303303
str(requested_fields))
304-
for t_name in requested_types:
304+
for t_name in requested_types: # pyrefly: ignore[not-iterable]
305305
t_class = getattr(types, t_name, None)
306-
if not issubclass(t_class, types.LitType):
306+
if not issubclass(t_class, types.LitType): # pyrefly: ignore[bad-argument-type]
307307
raise TypeError(f"Class '{t_name}' is not a valid LitType.")
308308
requested_fields.extend(utils.find_spec_keys(output_spec, t_class))
309-
ret_keys = set(requested_fields) # de-dupe
309+
ret_keys = set(requested_fields) # de-dupe # pyrefly: ignore[bad-argument-type]
310310

311311
# Return selected keys.
312312
logging.info('Will return keys: %s', str(ret_keys))
@@ -861,7 +861,7 @@ def _run_annotators(self,
861861
for annotator in self._annotators:
862862
annotator.annotate(datapoints, dataset, annotated_spec)
863863
return lit_dataset.Dataset(
864-
base=dataset, examples=datapoints, spec=annotated_spec)
864+
base=dataset, examples=datapoints, spec=annotated_spec) # pyrefly: ignore[bad-argument-type]
865865

866866
def make_handler(self, fn):
867867
"""Convenience wrapper to handle args and serialization.
@@ -1009,9 +1009,9 @@ def __init__(
10091009

10101010
# Interpreter initialization
10111011
if interpreters is not None:
1012-
self._interpreters = core.required_interpreters() | interpreters
1012+
self._interpreters = core.required_interpreters() | interpreters # pyrefly: ignore[unsupported-operation]
10131013
else:
1014-
self._interpreters = core.default_interpreters(self._models)
1014+
self._interpreters = core.default_interpreters(self._models) # pyrefly: ignore[bad-argument-type]
10151015

10161016
if metrics is not None:
10171017
self._metrics = metrics

0 commit comments

Comments
 (0)