diff --git a/fiddle/_src/absl_flags/flags.py b/fiddle/_src/absl_flags/flags.py index af5cb864..349c4810 100644 --- a/fiddle/_src/absl_flags/flags.py +++ b/fiddle/_src/absl_flags/flags.py @@ -264,9 +264,9 @@ def value(self): raise elif command == "set": - utils.set_value(self._value, expression) + utils.set_value(self._value, expression) # pyrefly: ignore[bad-argument-type] elif command == "fiddler": - self._value = self._apply_fiddler(self._value, expression) + self._value = self._apply_fiddler(self._value, expression) # pyrefly: ignore[bad-argument-type] else: raise AssertionError("Internal error; should not be reached.") return self._value diff --git a/fiddle/_src/absl_flags/legacy_flags.py b/fiddle/_src/absl_flags/legacy_flags.py index b5452e3b..d2bcba6d 100644 --- a/fiddle/_src/absl_flags/legacy_flags.py +++ b/fiddle/_src/absl_flags/legacy_flags.py @@ -214,7 +214,7 @@ def _rewrite(arg: str) -> str: explicit_name = 'fdl_tags_set' _, arg = arg.split('.', maxsplit=1) # Strip --fdl. or --fdl_tag. prefix. path, value = arg.split('=', maxsplit=1) - rewritten = f'--{explicit_name}={path}={value}' + rewritten = f'--{explicit_name}={path}={value}' # pyrefly: ignore[unbound-name] logging.debug('Rewrote flag "%s" to "%s".', arg, rewritten) return rewritten else: diff --git a/fiddle/_src/absl_flags/sweep_flag.py b/fiddle/_src/absl_flags/sweep_flag.py index 2cea591b..74a41cdf 100644 --- a/fiddle/_src/absl_flags/sweep_flag.py +++ b/fiddle/_src/absl_flags/sweep_flag.py @@ -162,7 +162,7 @@ def __init__( @property def value(self) -> Sequence[SweepItem]: if self._value is None: - self._value = self._parse(self._multi_flag.value) + self._value = self._parse(self._multi_flag.value) # pyrefly: ignore[bad-argument-type] return self._value def _parse_call_expression( diff --git a/fiddle/_src/absl_flags/utils_test.py b/fiddle/_src/absl_flags/utils_test.py index ae846666..8aa1a24a 100644 --- a/fiddle/_src/absl_flags/utils_test.py +++ b/fiddle/_src/absl_flags/utils_test.py @@ -201,7 +201,7 @@ def test_dotted_module_prefix_matching(self): class C: d = 42 - sub_b.c = C + sub_b.c = C # pyrefly: ignore[missing-attribute] sys.modules['a'] = parent_a sys.modules['a.b'] = sub_b diff --git a/fiddle/_src/codegen/auto_config/code_ir.py b/fiddle/_src/codegen/auto_config/code_ir.py index e7581b7f..40fdbd72 100644 --- a/fiddle/_src/codegen/auto_config/code_ir.py +++ b/fiddle/_src/codegen/auto_config/code_ir.py @@ -259,7 +259,7 @@ def to_stack(self) -> List[CallInstance]: result = [self] while current.parent is not None: current = current.parent - result.append(current) + result.append(current) # pyrefly: ignore[bad-argument-type] return list(reversed(result)) @arg_factory.supply_defaults diff --git a/fiddle/_src/codegen/auto_config/experimental_top_level_api.py b/fiddle/_src/codegen/auto_config/experimental_top_level_api.py index ada915e9..52d597f5 100644 --- a/fiddle/_src/codegen/auto_config/experimental_top_level_api.py +++ b/fiddle/_src/codegen/auto_config/experimental_top_level_api.py @@ -206,7 +206,7 @@ def code_generator( ImportSymbols(), TransformSubFixtures(), ] - passes.extend([ + passes.extend([ # pyrefly: ignore[bad-argument-type] LowerArgFactories(), MoveSharedNodesToVariables(), ]) @@ -214,17 +214,17 @@ def code_generator( is_complex = complex_to_variables.more_complex_than( max_expression_complexity ) - passes.append(MoveComplexNodesToVariables(is_complex=is_complex)) + passes.append(MoveComplexNodesToVariables(is_complex=is_complex)) # pyrefly: ignore[bad-argument-type] format_history = ( get_history_comments.format_history_for_buildable if include_history else make_symbolic_references.noop_history_comments ) - passes.extend([ + passes.extend([ # pyrefly: ignore[bad-argument-type] MakeSymbolicReferences(format_history=format_history), IrToCst(), ]) - return Codegen(passes=passes, debug_print=debug_print) + return Codegen(passes=passes, debug_print=debug_print) # pyrefly: ignore[bad-argument-type] def auto_config_codegen( diff --git a/fiddle/_src/codegen/auto_config/ir_printer.py b/fiddle/_src/codegen/auto_config/ir_printer.py index 5c732f3f..aa34d937 100644 --- a/fiddle/_src/codegen/auto_config/ir_printer.py +++ b/fiddle/_src/codegen/auto_config/ir_printer.py @@ -53,7 +53,7 @@ def _is_typing_or_generic(value: Any) -> bool: def format_py_reference(value: Any) -> str: - module_name = inspect.getmodule(value).__name__ + module_name = inspect.getmodule(value).__name__ # pyrefly: ignore[missing-attribute] if module_name in ("fiddle._src.config", "fiddle._src.partial"): module_name = "fdl" cls_name = value.__qualname__ diff --git a/fiddle/_src/codegen/auto_config/ir_to_cst.py b/fiddle/_src/codegen/auto_config/ir_to_cst.py index 8866a1b6..22388a01 100644 --- a/fiddle/_src/codegen/auto_config/ir_to_cst.py +++ b/fiddle/_src/codegen/auto_config/ir_to_cst.py @@ -114,7 +114,7 @@ def _prepare_args_helper( " please replace these objects in your input config, likely " "with fdl.Config nodes." ) - elements.append(cst.Element(sub_value)) + elements.append(cst.Element(sub_value)) # pyrefly: ignore[bad-argument-type] return cst_cls(elements) elif isinstance(value, dict): elements = [] @@ -130,9 +130,9 @@ def _prepare_args_helper( return cst.Attribute(value=base, attr=cst.Name(value.attribute)) elif isinstance(value, code_ir.ParameterizedTypeExpression): return cst.Subscript( - value=code_for_expr(value.base_expression), + value=code_for_expr(value.base_expression), # pyrefly: ignore[bad-argument-type] slice=[ - cst.SubscriptElement(cst.Index(code_for_expr(param))) + cst.SubscriptElement(cst.Index(code_for_expr(param))) # pyrefly: ignore[bad-argument-type] for param in value.param_expressions ], ) @@ -151,7 +151,7 @@ def _prepare_args_helper( ) args.extend( _prepare_args_helper( - names, values, attr, history=value.history_comments + names, values, attr, history=value.history_comments # pyrefly: ignore[bad-argument-type] ) ) if any_args_have_history: @@ -223,20 +223,20 @@ def code_for_fn( name = variable_decl.name.value assign = cst.Assign( targets=[cst.AssignTarget(target=cst.Name(name))], - value=code_for_expr(variable_decl.expression), + value=code_for_expr(variable_decl.expression), # pyrefly: ignore[bad-argument-type] ) variable_lines.append(cst.SimpleStatementLine(body=[assign])) body = cst.IndentedBlock( body=[ *variable_lines, cst.SimpleStatementLine( - body=[cst.Return(code_for_expr(fn.output_value))] + body=[cst.Return(code_for_expr(fn.output_value))] # pyrefly: ignore[bad-argument-type] ), ] ) if fn.return_type_annotation: returns = cst.Annotation( - annotation=code_for_expr(fn.return_type_annotation) + annotation=code_for_expr(fn.return_type_annotation) # pyrefly: ignore[bad-argument-type] ) else: returns = None diff --git a/fiddle/_src/codegen/auto_config/make_symbolic_references.py b/fiddle/_src/codegen/auto_config/make_symbolic_references.py index a0151d17..b057d71b 100644 --- a/fiddle/_src/codegen/auto_config/make_symbolic_references.py +++ b/fiddle/_src/codegen/auto_config/make_symbolic_references.py @@ -105,10 +105,10 @@ def _handle_partial( for name, arg_value in arguments.items(): if isinstance(arg_value, code_ir.ArgFactoryExpr): arg_factory_args[name] = state.call( - arg_value.expression, daglish.Attr(name) + arg_value.expression, daglish.Attr(name) # pyrefly: ignore[bad-argument-type] ) else: - regular_args[name] = state.call(arg_value, daglish.Attr(name)) + regular_args[name] = state.call(arg_value, daglish.Attr(name)) # pyrefly: ignore[bad-argument-type] for dict_of_args in (arg_factory_args, regular_args): for arg in dict_of_args: @@ -179,7 +179,7 @@ def traverse(value, state: daglish.State): return code_ir.SymbolOrFixtureCall( symbol_expression=ir_for_symbol, positional_arg_expressions=[], - arg_expressions=config_lib.ordered_arguments(value), + arg_expressions=config_lib.ordered_arguments(value), # pyrefly: ignore[bad-argument-type] history_comments=format_history(value), ) elif isinstance(value, partial.Partial): diff --git a/fiddle/_src/codegen/auto_config/sub_fixture.py b/fiddle/_src/codegen/auto_config/sub_fixture.py index 99f8f0d6..d7976608 100644 --- a/fiddle/_src/codegen/auto_config/sub_fixture.py +++ b/fiddle/_src/codegen/auto_config/sub_fixture.py @@ -73,13 +73,13 @@ def _is_super_ancestor( Returns: A bool value that indicates if `ancestor` is a super ancestor of all nodes. """ - nodes = list(nodes) + nodes = list(nodes) # pyrefly: ignore[bad-assignment] if len(nodes) == 0: # pylint: disable=g-explicit-length-test return False if len(nodes) == 1: if ancestor in nodes: return True - all_parents = node_to_parents_by_id[nodes[0]] + all_parents = node_to_parents_by_id[nodes[0]] # pyrefly: ignore[bad-index] if len(all_parents) == 1 and ancestor in all_parents: return True return _is_super_ancestor(ancestor, all_parents, node_to_parents_by_id) @@ -98,11 +98,11 @@ def _find_least_common_ancestor( ) -> int: """Find the least common ancestor of all nodes.""" - node_ids = list(node_ids) + node_ids = list(node_ids) # pyrefly: ignore[bad-assignment] if len(node_ids) == 0: # pylint: disable=g-explicit-length-test raise ValueError("Input nodes must not be empty.") if len(node_ids) == 1: - return node_ids[0] + return node_ids[0] # pyrefly: ignore[bad-index] if len(node_ids) == 2: x, y = node_ids if _is_super_ancestor(x, {y}, node_to_parents_by_id): @@ -114,8 +114,8 @@ def _find_least_common_ancestor( return _find_least_common_ancestor( x_parents.union(y_parents), node_to_parents_by_id ) - first_two = set(node_ids[:2]) - rest = set(node_ids[2:]) + first_two = set(node_ids[:2]) # pyrefly: ignore[bad-index] + rest = set(node_ids[2:]) # pyrefly: ignore[bad-index] first_two_ancestor = _find_least_common_ancestor( first_two, node_to_parents_by_id ) diff --git a/fiddle/_src/codegen/codegen_diff.py b/fiddle/_src/codegen/codegen_diff.py index 7cad0c59..a0534d8e 100644 --- a/fiddle/_src/codegen/codegen_diff.py +++ b/fiddle/_src/codegen/codegen_diff.py @@ -220,7 +220,7 @@ def _cst_for_fiddler(func_name: str, param_name: str, body: List[cst.CSTNode], name=cst.Name(func_name), params=cst.Parameters( params=[cst.Param(name=cst.Name(param_name), star='')]), - body=cst.IndentedBlock(body), + body=cst.IndentedBlock(body), # pyrefly: ignore[bad-argument-type] leading_lines=[cst.EmptyLine()] if add_leading_blank_line else []) @@ -237,7 +237,7 @@ def _cst_for_new_shared_value_variables( statements.append( cst.Assign( targets=[cst.AssignTarget(target=cst.Name(name))], - value=pyval_to_cst(value))) + value=pyval_to_cst(value))) # pyrefly: ignore[bad-argument-type] return [cst.SimpleStatementLine([stmt]) for stmt in statements] @@ -368,24 +368,24 @@ def _cst_for_changes(diff: diffing.Diff, param_name: str, new_value_cst = pyval_to_cst(change.new_value) update_callable = cst.Expr( cst.Call( - func=pyval_to_cst(mutate_buildable.update_callable), - args=[cst.Arg(parent_cst), cst.Arg(new_value_cst)], + func=pyval_to_cst(mutate_buildable.update_callable), # pyrefly: ignore[bad-argument-type] + args=[cst.Arg(parent_cst), cst.Arg(new_value_cst)], # pyrefly: ignore[bad-argument-type] ) ) elif isinstance(change, diffing.DeleteValue): - deletes.append(cst.Del(target=child_cst)) + deletes.append(cst.Del(target=child_cst)) # pyrefly: ignore[bad-argument-type] elif isinstance(change, diffing.RemoveTag): - arg_name = change.target[-1].name + arg_name = change.target[-1].name # pyrefly: ignore[missing-attribute] deletes.append( cst.Expr( cst.Call( - func=pyval_to_cst(tagging.remove_tag), + func=pyval_to_cst(tagging.remove_tag), # pyrefly: ignore[bad-argument-type] args=[ cst.Arg(parent_cst), - cst.Arg(pyval_to_cst(arg_name)), - cst.Arg(pyval_to_cst(change.tag)), + cst.Arg(pyval_to_cst(arg_name)), # pyrefly: ignore[bad-argument-type] + cst.Arg(pyval_to_cst(change.tag)), # pyrefly: ignore[bad-argument-type] ], ) ) @@ -395,18 +395,18 @@ def _cst_for_changes(diff: diffing.Diff, param_name: str, new_value_cst = pyval_to_cst(change.new_value) assigns.append( cst.Assign( - targets=[cst.AssignTarget(child_cst)], value=new_value_cst)) + targets=[cst.AssignTarget(child_cst)], value=new_value_cst)) # pyrefly: ignore[bad-argument-type] elif isinstance(change, diffing.AddTag): - arg_name = change.target[-1].name + arg_name = change.target[-1].name # pyrefly: ignore[missing-attribute] assigns.append( cst.Expr( value=cst.Call( - func=pyval_to_cst(tagging.add_tag), + func=pyval_to_cst(tagging.add_tag), # pyrefly: ignore[bad-argument-type] args=[ cst.Arg(parent_cst), - cst.Arg(pyval_to_cst(arg_name)), - cst.Arg(pyval_to_cst(change.tag)), + cst.Arg(pyval_to_cst(arg_name)), # pyrefly: ignore[bad-argument-type] + cst.Arg(pyval_to_cst(change.tag)), # pyrefly: ignore[bad-argument-type] ], ) ) @@ -433,17 +433,17 @@ def _cst_for_child(parent_cst: cst.CSTNode, child_path_elt: daglish.PathElement, pyval_to_cst: A function used to convert Python values to CST. """ if isinstance(child_path_elt, daglish.Attr): - return cst.Attribute(value=parent_cst, attr=cst.Name(child_path_elt.name)) + return cst.Attribute(value=parent_cst, attr=cst.Name(child_path_elt.name)) # pyrefly: ignore[bad-argument-type] elif isinstance(child_path_elt, daglish.Index): index_cst = pyval_to_cst(child_path_elt.index) return cst.Subscript( - value=parent_cst, - slice=[cst.SubscriptElement(slice=cst.Index(index_cst))]) + value=parent_cst, # pyrefly: ignore[bad-argument-type] + slice=[cst.SubscriptElement(slice=cst.Index(index_cst))]) # pyrefly: ignore[bad-argument-type] elif isinstance(child_path_elt, daglish.Key): key_cst = pyval_to_cst(child_path_elt.key) return cst.Subscript( - value=parent_cst, - slice=[cst.SubscriptElement(slice=cst.Index(key_cst))]) + value=parent_cst, # pyrefly: ignore[bad-argument-type] + slice=[cst.SubscriptElement(slice=cst.Index(key_cst))]) # pyrefly: ignore[bad-argument-type] else: raise ValueError(f'Unsupported PathElement {type(child_path_elt)}') diff --git a/fiddle/_src/codegen/import_manager.py b/fiddle/_src/codegen/import_manager.py index 6902fb4a..d003393c 100644 --- a/fiddle/_src/codegen/import_manager.py +++ b/fiddle/_src/codegen/import_manager.py @@ -50,10 +50,10 @@ def parse_import(stmt: str) -> AnyImport: def _get_import_name_node(node: AnyImport) -> cst.ImportAlias: - if len(node.names) != 1: + if len(node.names) != 1: # pyrefly: ignore[bad-argument-type] raise ValueError( f"CST nodes with more than 1 name are not supported; got {node}") - return node.names[0] + return node.names[0] # pyrefly: ignore[bad-index] def get_import_name(node: AnyImport) -> str: @@ -95,7 +95,7 @@ def get_full_module_name(node: AnyImport) -> str: name_str = _dummy_module_for_formatting.code_for_node( _get_import_name_node(node).name) if isinstance(node, cst.ImportFrom): - module_str = _dummy_module_for_formatting.code_for_node(node.module) + module_str = _dummy_module_for_formatting.code_for_node(node.module) # pyrefly: ignore[bad-argument-type] return f"{module_str}.{name_str}" else: return name_str @@ -230,7 +230,7 @@ def add(self, value: Any) -> str: Returns: Relative-qualified name for the instance. """ - module_name = inspect.getmodule(value).__name__ + module_name = inspect.getmodule(value).__name__ # pyrefly: ignore[missing-attribute] if isinstance(value, enum.Enum): value_qualname = value.__class__.__qualname__ + "." + value.name else: diff --git a/fiddle/_src/codegen/legacy_codegen.py b/fiddle/_src/codegen/legacy_codegen.py index 11e651c7..064bebe7 100644 --- a/fiddle/_src/codegen/legacy_codegen.py +++ b/fiddle/_src/codegen/legacy_codegen.py @@ -164,7 +164,7 @@ def traverse(child, state=None): return state.map_children(child) else: return special_value_codegen.transform_py_value(child, - self.import_manager) + self.import_manager) # pyrefly: ignore[bad-argument-type] lhs = assignment_path(lhs_var, lhs_path) assignment = mini_ast.Assignment(lhs, repr(traverse(attr_value))) @@ -208,12 +208,12 @@ def traverse(child, state: daglish.State): mini_ast.Assignment(name, f"{buildable_subclass_str}({relname})") ] for key, value in child.__arguments__.items(): - path = [daglish.BuildableAttr(key)] - nodes.append(shared_manager.assign(name, path, value)) + path = [daglish.BuildableAttr(key)] # pyrefly: ignore[bad-argument-type] + nodes.append(shared_manager.assign(name, path, value)) # pyrefly: ignore[bad-argument-type] # `shared_manager` indexes by ID, so be careful to use the original DAG # node `child`. - shared_manager.add(name, child, mini_ast.ImmediateAttrsBlock(nodes)) + shared_manager.add(name, child, mini_ast.ImmediateAttrsBlock(nodes)) # pyrefly: ignore[bad-argument-type] traverser = daglish.MemoizedTraversal( traverse, @@ -290,14 +290,14 @@ def handle_child_attr(value, state: daglish.State): config_lib.Buildable) and value not in shared_manager: deferred.append((value, full_path)) else: - nodes.append(shared_manager.assign("root", full_path, value)) + nodes.append(shared_manager.assign("root", full_path, value)) # pyrefly: ignore[bad-argument-type] if state.is_traversable(value) and not isinstance(value, config_lib.Buildable): state.flattened_map_children(value) daglish.BasicTraversal.run(handle_child_attr, child) - main_tree_blocks.append(mini_ast.ImmediateAttrsBlock(nodes)) + main_tree_blocks.append(mini_ast.ImmediateAttrsBlock(nodes)) # pyrefly: ignore[bad-argument-type] # Recurses to configure sub-Buildable nodes. for sub_child, sub_path in deferred: diff --git a/fiddle/_src/codegen/mini_ast.py b/fiddle/_src/codegen/mini_ast.py index 79f96158..193a6350 100644 --- a/fiddle/_src/codegen/mini_ast.py +++ b/fiddle/_src/codegen/mini_ast.py @@ -93,7 +93,7 @@ def block(sub_nodes_or_lines: List[Union[str, List[str], CodegenNode]], Returns: List of lines, taken from `sub_nodes_or_lines`. """ - sub_nodes_or_lines = filter(None, sub_nodes_or_lines) # Remove empty items. + sub_nodes_or_lines = filter(None, sub_nodes_or_lines) # Remove empty items. # pyrefly: ignore[bad-assignment] result = [] for i, item in enumerate(sub_nodes_or_lines): @@ -117,7 +117,7 @@ class ImmediateAttrsBlock(CodegenNode): nodes: List[CodegenNode] def lines(self) -> List[str]: - return block(self.nodes, []) + return block(self.nodes, []) # pyrefly: ignore[bad-argument-type] @dataclasses.dataclass(frozen=True) @@ -127,7 +127,7 @@ class SharedThenResultAssignment(CodegenNode): tree_blocks: List[CodegenNode] def lines(self) -> List[str]: - return block(self.shared_instances + self.tree_blocks, [""]) + return block(self.shared_instances + self.tree_blocks, [""]) # pyrefly: ignore[bad-argument-type] @dataclasses.dataclass(frozen=True) @@ -141,6 +141,6 @@ class ConfigBuilder(CodegenNode): def lines(self) -> List[str]: import_lines = cst.Module(body=self.imports).code.splitlines() - fn_body = [" " + line for line in block(self.builder_body, [""])] + fn_body = [" " + line for line in block(self.builder_body, [""])] # pyrefly: ignore[bad-argument-type] return block([block([import_lines], []), ["def build_config():"] + fn_body], ["", ""]) diff --git a/fiddle/_src/codegen/new_codegen.py b/fiddle/_src/codegen/new_codegen.py index 108ef4a4..801a9a95 100644 --- a/fiddle/_src/codegen/new_codegen.py +++ b/fiddle/_src/codegen/new_codegen.py @@ -74,7 +74,7 @@ def _get_pass_idx( cls: Type[experimental_top_level_api.CodegenPass], ) -> int: for i, codegen_pass in enumerate(codegen_config.passes): - if issubclass(fdl.get_callable(codegen_pass), cls): + if issubclass(fdl.get_callable(codegen_pass), cls): # pyrefly: ignore[bad-argument-type] return i raise ValueError(f"Could not find codegen pass {cls}") diff --git a/fiddle/_src/codegen/newcg_symbolic_references.py b/fiddle/_src/codegen/newcg_symbolic_references.py index 031f40b8..e4e091a9 100644 --- a/fiddle/_src/codegen/newcg_symbolic_references.py +++ b/fiddle/_src/codegen/newcg_symbolic_references.py @@ -99,7 +99,7 @@ def traverse(value, state: daglish.State): return code_ir.SymbolOrFixtureCall( symbol_expression=ir_for_buildable_type, positional_arg_expressions=[ir_for_symbol], - arg_expressions=config_lib.ordered_arguments(value), + arg_expressions=config_lib.ordered_arguments(value), # pyrefly: ignore[bad-argument-type] history_comments=format_history(value), ) elif is_plain_symbol_or_enum_value(value): diff --git a/fiddle/_src/codegen/py_val_to_cst_converter.py b/fiddle/_src/codegen/py_val_to_cst_converter.py index b9a0139d..05f45a3b 100644 --- a/fiddle/_src/codegen/py_val_to_cst_converter.py +++ b/fiddle/_src/codegen/py_val_to_cst_converter.py @@ -152,10 +152,10 @@ def register_py_val_to_cst_converter(matchers: Union[ValueMatcher, A decorator function. """ if not isinstance(matchers, list): - matchers = [matchers] + matchers = [matchers] # pyrefly: ignore[bad-assignment] def decorator(converter: ValueConverterFunc) -> ValueConverterFunc: - for matcher in matchers: + for matcher in matchers: # pyrefly: ignore[not-iterable] if priority is None: matcher_priority = 100 if isinstance(matcher, type) else 50 else: @@ -284,20 +284,20 @@ def _convert_complex(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: @register_py_val_to_cst_converter(list) def _convert_list(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts a list to CST.""" - return cst.List([cst.Element(conversion_fn(v)) for v in value]) + return cst.List([cst.Element(conversion_fn(v)) for v in value]) # pyrefly: ignore[bad-argument-type] @register_py_val_to_cst_converter(tuple) def _convert_tuple(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts a tuple to CST.""" - return cst.Tuple([cst.Element(conversion_fn(v)) for v in value]) + return cst.Tuple([cst.Element(conversion_fn(v)) for v in value]) # pyrefly: ignore[bad-argument-type] @register_py_val_to_cst_converter(dict) def _convert_dict(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts a dict to CST.""" return cst.Dict([ - cst.DictElement(conversion_fn(key), conversion_fn(val)) + cst.DictElement(conversion_fn(key), conversion_fn(val)) # pyrefly: ignore[bad-argument-type] for (key, val) in value.items() ]) @@ -306,7 +306,7 @@ def _convert_dict(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: def _convert_set(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts a set to CST.""" if value: - return cst.Set([cst.Element(conversion_fn(v)) for v in value]) + return cst.Set([cst.Element(conversion_fn(v)) for v in value]) # pyrefly: ignore[bad-argument-type] else: return cst.Call(func=cst.Name('set')) @@ -317,9 +317,9 @@ def _convert_slice(value: slice, conversion_fn: PyValToCstFunc) -> cst.CSTNode: return cst.Call( func=cst.Name('slice'), args=[ - cst.Arg(conversion_fn(value.start)), - cst.Arg(conversion_fn(value.stop)), - cst.Arg(conversion_fn(value.step)), + cst.Arg(conversion_fn(value.start)), # pyrefly: ignore[bad-argument-type] + cst.Arg(conversion_fn(value.stop)), # pyrefly: ignore[bad-argument-type] + cst.Arg(conversion_fn(value.step)), # pyrefly: ignore[bad-argument-type] ], ) @@ -329,7 +329,7 @@ def _convert_namedtuple(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts an instance of a named tuple to CST.""" return cst.Call( - func=conversion_fn(type(value)), + func=conversion_fn(type(value)), # pyrefly: ignore[bad-argument-type] args=[ kwarg_to_cst(arg_name, conversion_fn(arg_val)) for (arg_name, arg_val) in value._asdict().items() @@ -340,13 +340,13 @@ def _convert_namedtuple(value: Any, def _convert_buildable(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts a fdl.Config or fdl.Partial to CST.""" - args = [cst.Arg(conversion_fn(config_lib.get_callable(value)))] + args = [cst.Arg(conversion_fn(config_lib.get_callable(value)))] # pyrefly: ignore[bad-argument-type] for (arg_name, arg_val) in value.__arguments__.items(): if arg_name in value.__argument_tags__: for tag in value.__argument_tags__[arg_name]: arg_val = tag.new(arg_val) args.append(kwarg_to_cst(arg_name, conversion_fn(arg_val))) - return cst.Call(func=conversion_fn(type(value)), args=args) + return cst.Call(func=conversion_fn(type(value)), args=args) # pyrefly: ignore[bad-argument-type] @register_py_val_to_cst_converter(tagging.TaggedValueCls) @@ -356,8 +356,8 @@ def _convert_tagged_value(value: Any, node = conversion_fn(value.value) for tag in sorted(value.tags, key=repr, reverse=True): node = cst.Call( - func=cst.Attribute(value=conversion_fn(tag), attr=cst.Name('new')), - args=[cst.Arg(node)]) + func=cst.Attribute(value=conversion_fn(tag), attr=cst.Name('new')), # pyrefly: ignore[bad-argument-type] + args=[cst.Arg(node)]) # pyrefly: ignore[bad-argument-type] return node @@ -391,12 +391,12 @@ def _convert_importable(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts an importable value to the CST for `.`.""" module = inspect.getmodule(value) - if module.__name__ == '__main__' or module is builtins: + if module.__name__ == '__main__' or module is builtins: # pyrefly: ignore[missing-attribute] return dotted_name_to_cst(value.__qualname__) else: result = conversion_fn(inspect.getmodule(value)) for piece in value.__qualname__.split('.'): - result = cst.Attribute(value=result, attr=cst.Name(piece)) + result = cst.Attribute(value=result, attr=cst.Name(piece)) # pyrefly: ignore[bad-argument-type] return result @@ -405,9 +405,9 @@ def _convert_partial(value: functools.partial, conversion_fn: PyValToCstFunc) -> cst.CSTNode: """Converts a functools.partial to CST.""" return cst.Call( - func=conversion_fn(functools.partial), - args=([cst.Arg(conversion_fn(value.func))] + - [cst.Arg(conversion_fn(arg)) for arg in value.args] + [ + func=conversion_fn(functools.partial), # pyrefly: ignore[bad-argument-type] + args=([cst.Arg(conversion_fn(value.func))] + # pyrefly: ignore[bad-argument-type] + [cst.Arg(conversion_fn(arg)) for arg in value.args] + [ # pyrefly: ignore[bad-argument-type] kwarg_to_cst(arg_name, conversion_fn(arg_val)) for (arg_name, arg_val) in value.keywords.items() ])) @@ -416,5 +416,5 @@ def _convert_partial(value: functools.partial, @register_py_val_to_cst_converter(lambda value: isinstance(value, enum.Enum)) def _convert_enum(value: Any, conversion_fn: PyValToCstFunc) -> cst.CSTNode: return cst.Attribute( - value=conversion_fn(type(value)), attr=cst.Name(value.name) + value=conversion_fn(type(value)), attr=cst.Name(value.name) # pyrefly: ignore[bad-argument-type] ) diff --git a/fiddle/_src/codegen/py_val_to_cst_converter_test.py b/fiddle/_src/codegen/py_val_to_cst_converter_test.py index e79e96ae..7606dcb3 100644 --- a/fiddle/_src/codegen/py_val_to_cst_converter_test.py +++ b/fiddle/_src/codegen/py_val_to_cst_converter_test.py @@ -130,13 +130,13 @@ class PyValToCstConverterTest(parameterized.TestCase): ]) def test_convert(self, pyval, expected): cst_expr = py_val_to_cst_converter.convert_py_val_to_cst(pyval) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual(_get_cst_code(cst_module), expected + '\n') def test_convert_multiple_tags(self): pyval = fdl.TaggedValue([SampleTag, AnotherTag], 3) cst_expr = py_val_to_cst_converter.convert_py_val_to_cst(pyval) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual( _get_cst_code(cst_module), 'AnotherTag.new(SampleTag.new(3))\n') @@ -144,7 +144,7 @@ def test_convert_new_tags(self): pyval = fdl.Config(SampleNamedTuple, x=1) fdl.add_tag(pyval, 'x', SampleTag) cst_expr = py_val_to_cst_converter.convert_py_val_to_cst(pyval) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual( _get_cst_code(cst_module), 'fiddle._src.config.Config(SampleNamedTuple, x=SampleTag.new(1))\n', @@ -152,7 +152,7 @@ def test_convert_new_tags(self): def test_convert_empty_set(self): cst_expr = py_val_to_cst_converter.convert_py_val_to_cst(set()) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual(_get_cst_code(cst_module), 'set()\n') def test_convert_unsupported_type(self): @@ -179,13 +179,13 @@ def convert_named_value(value, convert_child, id_to_name): cst_expr = py_val_to_cst_converter.convert_py_val_to_cst( pyval, [custom_converter]) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual( _get_cst_code(cst_module), "[1, {2: x}, MyFiddleConfig(re.match, pattern='a|b')]\n") cst_expr = py_val_to_cst_converter.convert_py_val_to_cst(pyval) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual( _get_cst_code(cst_module), "[1, {2: [1]}, fiddle._src.config.Config(re.match, pattern='a|b')]\n", diff --git a/fiddle/_src/experimental/auto_config.py b/fiddle/_src/experimental/auto_config.py index 2d7df69d..2f4419b4 100644 --- a/fiddle/_src/experimental/auto_config.py +++ b/fiddle/_src/experimental/auto_config.py @@ -345,7 +345,7 @@ def visit_GeneratorExp(self, node: ast.GeneratorExp): def visit_Try(self, node: ast.Try): return self._handle_control_flow(node) - def visit_Raise(self, node: ast.Try): + def visit_Raise(self, node: ast.Try): # pyrefly: ignore[bad-override] return self._handle_control_flow(node, activatable=True) def visit_With(self, node: ast.With): @@ -477,13 +477,13 @@ def fn(...): # Or some expression involving a lambda. def _find_function_code(code: types.CodeType, fn_name: str): """Finds the code object within `code` corresponding to `fn_name`.""" - code = [ + code = [ # pyrefly: ignore[bad-assignment] const for const in code.co_consts if inspect.iscode(const) and const.co_name == fn_name ] - assert len(code) == 1, f"Couldn't find function code for {fn_name!r}." - return code[0] + assert len(code) == 1, f"Couldn't find function code for {fn_name!r}." # pyrefly: ignore[bad-argument-type] + return code[0] # pyrefly: ignore[bad-index] def _unwrap_code_for_fn(code: types.CodeType, fn: types.FunctionType): @@ -512,7 +512,7 @@ def _make_closure_cell(contents): else: # For earlier versions of Python, build a dummy function to get CellType. dummy_fn = lambda: contents - cell_type = type(dummy_fn.__closure__[0]) + cell_type = type(dummy_fn.__closure__[0]) # pyrefly: ignore[unsupported-operation] return cell_type(contents) @@ -615,7 +615,7 @@ def build_model(): A wrapped version of the same callable that will not be transformed to config if called inside an auto_config function. """ - return AutoConfig( + return AutoConfig( # pyrefly: ignore[bad-return] func=fn_or_cls, buildable_func=fn_or_cls, always_inline=True ) @@ -902,7 +902,7 @@ def make_auto_config(fn): line_number = fn.__code__.co_firstlineno node_transformer = _AutoConfigNodeTransformer( source=source, - filename=filename, + filename=filename, # pyrefly: ignore[bad-argument-type] line_number=line_number, allow_control_flow=experimental_allow_control_flow, ) @@ -924,7 +924,7 @@ def make_auto_config(fn): node = _wrap_ast_for_fn_with_closure_vars(node, fn) # Compile the modified AST, and then find the function code object within # the returned module-level code object. - code = compile(node, inspect.getsourcefile(fn), 'exec') + code = compile(node, inspect.getsourcefile(fn), 'exec') # pyrefly: ignore[bad-argument-type] code = _unwrap_code_for_fn(code, fn) # Insert auto_config_attr_load_handler, auto_config_attr_save_handler, @@ -988,7 +988,7 @@ def as_buildable(*args, **kwargs): fn = method_type(fn) as_buildable = method_type(as_buildable) return AutoConfig( - fn, as_buildable, always_inline=experimental_always_inline + fn, as_buildable, always_inline=experimental_always_inline # pyrefly: ignore[bad-argument-type] ) # Decorator with empty parenthesis. @@ -1142,7 +1142,7 @@ def make_experiment(): ) # Evaluate the `as_buildable` interpretation. auto_config_fn = cast(AutoConfig, buildable.__fn_or_cls__) - tmp_config = auto_config_fn.as_buildable(**buildable.__arguments__) + tmp_config = auto_config_fn.as_buildable(**buildable.__arguments__) # pyrefly: ignore[bad-unpacking] if not isinstance(tmp_config, config.Buildable): raise ValueError( 'You cannot currently inline functions that do not return ' @@ -1187,7 +1187,7 @@ def __init__(self, lambda_fn): def visit_Lambda(self, node) -> None: loc = self.get_metadata(cst.metadata.PositionProvider, node) - if loc.start.line == self.lineno: + if loc.start.line == self.lineno: # pyrefly: ignore[missing-attribute] self.candidates.append(node) @@ -1196,7 +1196,7 @@ def _getsource_for_lambda(fn: Callable[..., Any]) -> str: # Get the source for the module that defines `fn`. module = inspect.getmodule(fn) filename = inspect.getsourcefile(fn) - lines = linecache.getlines(filename, module.__dict__) + lines = linecache.getlines(filename, module.__dict__) # pyrefly: ignore[bad-argument-type] source = ''.join(lines) # Parse the CST for the module, and search for the lambda. diff --git a/fiddle/_src/experimental/autobuilders/autobuilders_test.py b/fiddle/_src/experimental/autobuilders/autobuilders_test.py index 09400b6c..58bcce0b 100644 --- a/fiddle/_src/experimental/autobuilders/autobuilders_test.py +++ b/fiddle/_src/experimental/autobuilders/autobuilders_test.py @@ -54,7 +54,7 @@ class Foo: def __init__(self, x): self.x = x - @ab.skeleton(Foo) + @ab.skeleton(Foo) # pyrefly: ignore[bad-argument-type] def foo_skeleton(config: fdl.Config): # pylint: disable=unused-variable # Note: Setting constants is generally not the purpose of skeletons; we're # just doing that here for testing purposes. @@ -78,23 +78,23 @@ class Foo: # Still raises informative error even when table entry is present # (for a validator). - ab.validator(Foo)(lambda config: None) + ab.validator(Foo)(lambda config: None) # pyrefly: ignore[bad-argument-type] with self.assertRaisesRegex(KeyError, r".*\bFoo\b"): ab.config(Foo) def test_skeleton_registers(self): registry = ab.Registry() fn = lambda config: None - registry.skeleton(FakeDense)(fn) + registry.skeleton(FakeDense)(fn) # pyrefly: ignore[bad-argument-type] self.assertDictEqual(registry.table, { - FakeDense: ab.TableEntry(skeleton_fn=fn, validators=[]), + FakeDense: ab.TableEntry(skeleton_fn=fn, validators=[]), # pyrefly: ignore[bad-argument-type] }) def test_skeleton_duplicate_class_registration_error(self): registry = ab.Registry() - registry.skeleton(FakeDense)(lambda config: None) + registry.skeleton(FakeDense)(lambda config: None) # pyrefly: ignore[bad-argument-type] with self.assertRaisesRegex(ab.DuplicateSkeletonError, r".*FakeDense.*"): - registry.skeleton(FakeDense)(lambda config: None) + registry.skeleton(FakeDense)(lambda config: None) # pyrefly: ignore[bad-argument-type] def test_skeleton_duplicate_function_registration_error(self): # Because the fancy error message includes source lines, make sure that @@ -103,18 +103,18 @@ def fake_fn(): pass registry = ab.Registry() - registry.skeleton(fake_fn)(lambda config: None) + registry.skeleton(fake_fn)(lambda config: None) # pyrefly: ignore[bad-argument-type] with self.assertRaisesRegex(ab.DuplicateSkeletonError, r".*fake_fn.*"): - registry.skeleton(fake_fn)(lambda config: None) + registry.skeleton(fake_fn)(lambda config: None) # pyrefly: ignore[bad-argument-type] def test_recursive_skeleton(self): - @ab.skeleton(FakeDense) + @ab.skeleton(FakeDense) # pyrefly: ignore[bad-argument-type] def dense_skeleton(config: fdl.Config) -> None: # pylint: disable=unused-variable config.in_dim = 4 config.out_dim = 4 - @ab.skeleton(FakeMlp) + @ab.skeleton(FakeMlp) # pyrefly: ignore[bad-argument-type] def mlp_skeleton(config: fdl.Config) -> None: # pylint: disable=unused-variable config.first_dense = ab.config(FakeDense) config.first_dense.in_dim = 5 @@ -136,12 +136,12 @@ def test_auto_skeleton_basic(self): def test_auto_skeleton_subclasses_and_existing_skeletons(self): - @ab.skeleton(FakeDense) + @ab.skeleton(FakeDense) # pyrefly: ignore[bad-argument-type] def dense_skeleton(config: fdl.Config) -> None: # pylint: disable=unused-variable config.in_dim = 4 config.out_dim = 4 - @ab.skeleton(FakeDenseSubclass) + @ab.skeleton(FakeDenseSubclass) # pyrefly: ignore[bad-argument-type] def dense_subclass_skeleton(config: fdl.Config) -> None: # pylint: disable=unused-variable config.in_dim = 7 config.out_dim = 9 @@ -210,9 +210,9 @@ def foo(x): def test_validator_registers(self): registry = ab.Registry() fn = lambda config: None - registry.validator(FakeDense)(fn) + registry.validator(FakeDense)(fn) # pyrefly: ignore[bad-argument-type] self.assertDictEqual(registry.table, { - FakeDense: ab.TableEntry(skeleton_fn=None, validators=[fn]), + FakeDense: ab.TableEntry(skeleton_fn=None, validators=[fn]), # pyrefly: ignore[bad-argument-type] }) diff --git a/fiddle/_src/experimental/daglish_legacy.py b/fiddle/_src/experimental/daglish_legacy.py index 6961ee36..89cfd9cc 100644 --- a/fiddle/_src/experimental/daglish_legacy.py +++ b/fiddle/_src/experimental/daglish_legacy.py @@ -258,7 +258,7 @@ def wrap_with_paths(current_path: daglish.Path, value: Any): parent = daglish.follow_path(structure, current_path[:-1]) parent_paths = paths_memo[id(parent)] all_paths = daglish.add_path_element(parent_paths, current_path[-1]) - return (yield from fn(all_paths, current_path, value)) + return (yield from fn(all_paths, current_path, value)) # pyrefly: ignore[bad-argument-type] return traverse_with_path(wrap_with_paths, structure) diff --git a/fiddle/_src/experimental/dataclasses.py b/fiddle/_src/experimental/dataclasses.py index 92d6f71c..2d81be7a 100644 --- a/fiddle/_src/experimental/dataclasses.py +++ b/fiddle/_src/experimental/dataclasses.py @@ -54,7 +54,7 @@ class DataclassTraverserRegistry(daglish.NodeTraverserRegistry): def find_node_traverser( self, node_type: Type[Any] ) -> Optional[daglish.NodeTraverser]: - traverser = self.fallback_registry.find_node_traverser(node_type) + traverser = self.fallback_registry.find_node_traverser(node_type) # pyrefly: ignore[missing-attribute] if traverser is None and dataclasses.is_dataclass(node_type): traverser = dataclass_traverser return traverser diff --git a/fiddle/_src/experimental/lazy_imports_test.py b/fiddle/_src/experimental/lazy_imports_test.py index 4dd66660..736e0726 100644 --- a/fiddle/_src/experimental/lazy_imports_test.py +++ b/fiddle/_src/experimental/lazy_imports_test.py @@ -39,7 +39,7 @@ def sum(self) -> Any: return a + b -class LazyImportsInspectTest(absltest.TestCase, unittest.TestCase): +class LazyImportsInspectTest(absltest.TestCase, unittest.TestCase): # pyrefly: ignore[inconsistent-inheritance] def assert_module(self, m: lazy_imports.ProxyObject, name: str) -> None: self.assertIsInstance(m, lazy_imports.ProxyObject) @@ -71,11 +71,11 @@ def test_proxy_objects(self): # pylint: enable=g-import-not-at-top,g-multiple-import with self.subTest('qualname'): - self.assert_module(a0, 'a0') - self.assert_module(a1.b.c, 'a1.b.c') + self.assert_module(a0, 'a0') # pyrefly: ignore[bad-argument-type] + self.assert_module(a1.b.c, 'a1.b.c') # pyrefly: ignore[bad-argument-type] self.assert_module(a1.non_module.c, 'a1:non_module.c') - self.assert_module(c00, 'a2.b.c') - self.assert_module(c01, 'a2.b.c') + self.assert_module(c00, 'a2.b.c') # pyrefly: ignore[bad-argument-type] + self.assert_module(c01, 'a2.b.c') # pyrefly: ignore[bad-argument-type] self.assert_module(c02, 'a2.b.c') self.assert_module(c02.non_module.c, 'a2.b.c:non_module.c') self.assert_module(c2, 'a3.c2') @@ -89,36 +89,36 @@ def test_proxy_objects(self): self.assertIs(c02, c00) with self.subTest('inspect_signature'): - self.assert_signature(a0, kw_only=True) - self.assert_signature(a1.b.c, kw_only=True) - self.assert_signature(c00, kw_only=True) - self.assert_signature(c01, kw_only=True) - self.assert_signature(c02, kw_only=True) + self.assert_signature(a0, kw_only=True) # pyrefly: ignore[bad-argument-type] + self.assert_signature(a1.b.c, kw_only=True) # pyrefly: ignore[bad-argument-type] + self.assert_signature(c00, kw_only=True) # pyrefly: ignore[bad-argument-type] + self.assert_signature(c01, kw_only=True) # pyrefly: ignore[bad-argument-type] + self.assert_signature(c02, kw_only=True) # pyrefly: ignore[bad-argument-type] with self.subTest('inspect_getmodule'): self.assertEqual( - inspect.getmodule(a0).__name__, + inspect.getmodule(a0).__name__, # pyrefly: ignore[missing-attribute] 'fiddle._src.experimental.lazy_imports', ) self.assertEqual( - inspect.getmodule(a1.b.c).__name__, + inspect.getmodule(a1.b.c).__name__, # pyrefly: ignore[missing-attribute] 'fiddle._src.experimental.lazy_imports', ) self.assertEqual( - inspect.getmodule(c00).__name__, + inspect.getmodule(c00).__name__, # pyrefly: ignore[missing-attribute] 'fiddle._src.experimental.lazy_imports', ) self.assertEqual( - inspect.getmodule(c01).__name__, + inspect.getmodule(c01).__name__, # pyrefly: ignore[missing-attribute] 'fiddle._src.experimental.lazy_imports', ) self.assertEqual( - inspect.getmodule(c02).__name__, + inspect.getmodule(c02).__name__, # pyrefly: ignore[missing-attribute] 'fiddle._src.experimental.lazy_imports', ) -class BuildLazyImportsTest(absltest.TestCase, unittest.TestCase): +class BuildLazyImportsTest(absltest.TestCase, unittest.TestCase): # pyrefly: ignore[inconsistent-inheritance] def test_import_as(self): with lazy_imports.lazy_imports(kw_only=True): @@ -206,7 +206,7 @@ def test_kw_only_check(self): _ = config_lib.Config(lazy_imports_test_example.MyDataClass, 1, 2) -class SerializationTest(absltest.TestCase, unittest.TestCase): +class SerializationTest(absltest.TestCase, unittest.TestCase): # pyrefly: ignore[inconsistent-inheritance] def test_regular_import(self): from fiddle._src.experimental import lazy_imports_test_example # pylint: disable=g-import-not-at-top diff --git a/fiddle/_src/experimental/serialization.py b/fiddle/_src/experimental/serialization.py index 257e7f80..f37d09d7 100644 --- a/fiddle/_src/experimental/serialization.py +++ b/fiddle/_src/experimental/serialization.py @@ -166,7 +166,7 @@ def register_dict_based_object(object_type: Type[Any]): register_node_traverser( bytes, flatten_fn=lambda x: ((x.decode('raw_unicode_escape'),), None), - unflatten_fn=lambda values, _: values[0].encode('raw_unicode_escape'), + unflatten_fn=lambda values, _: values[0].encode('raw_unicode_escape'), # pyrefly: ignore[bad-index] path_elements_fn=lambda x: (IdentityElement(),), ) @@ -224,7 +224,7 @@ def allows_import(self, module: str, symbol: str) -> bool: symbol: The symbol to import from `module`. """ - def allows_value(self, value: Any) -> bool: + def allows_value(self, value: Any) -> bool: # pyrefly: ignore[bad-return] """Returns whether this policy allows an imported `value` to be used. This is called after `value` has already been imported, but before it is @@ -647,7 +647,7 @@ def _serialize( # If we should add paths (all_paths is not None) and we have an entry for # value in self._paths_by_id, use that, since it may contain additional # paths not available via the parent. - all_paths = self._paths_by_id.get(id(value), all_paths) + all_paths = self._paths_by_id.get(id(value), all_paths) # pyrefly: ignore[bad-assignment] traverser = find_node_traverser(type(value)) if traverser is None: diff --git a/fiddle/_src/experimental/visualize.py b/fiddle/_src/experimental/visualize.py index 08ae7dbc..aa5737d6 100644 --- a/fiddle/_src/experimental/visualize.py +++ b/fiddle/_src/experimental/visualize.py @@ -146,7 +146,7 @@ def traverse_fn(value, state: daglish.State): should_copy = True for name, attr_value in list(value.__arguments__.items()): - param = value.__signature_info__.parameters.get(name, None) + param = value.__signature_info__.parameters.get(name, None) # pyrefly: ignore[no-matching-overload] if param is None: continue param_default = ( @@ -158,12 +158,12 @@ def traverse_fn(value, state: daglish.State): # All paths must flow through both the parent config (`value`) and # the specific attribute which is being defaulted, in order for us # to safely remove it. - and can_remove_deep_default(attr_value, name, state) + and can_remove_deep_default(attr_value, name, state) # pyrefly: ignore[bad-argument-type] ): if should_copy: value = copy.copy(value) should_copy = False - delattr(value, name) + delattr(value, name) # pyrefly: ignore[bad-argument-type] return state.map_children(value) return daglish.MemoizedTraversal.run(traverse_fn, config) @@ -236,7 +236,7 @@ def traverse(value, state: daglish.State): result = state.map_children(value) for name, sub_value in config_lib.ordered_arguments(result).items(): if sub_value is _any_value: - delattr(result, name) + delattr(result, name) # pyrefly: ignore[bad-argument-type] return result else: result = state.flattened_map_children(value) @@ -288,7 +288,7 @@ def traverse(value, state: daglish.State): to_keep = fields_by_id[id(value)] value = copy.copy(value) # Shallow copy to avoid mutating original. for argument in set(config_lib.ordered_arguments(value)) - set(to_keep): - setattr(value, argument, Trimmed()) + setattr(value, argument, Trimmed()) # pyrefly: ignore[bad-argument-type] return state.map_children(value) return daglish.MemoizedTraversal.run(traverse, config) @@ -324,7 +324,7 @@ def trim_long_fields( def traverse(value, state: daglish.State): if isinstance(value, config_lib.Buildable): for argument in set(config_lib.ordered_arguments(value)): - field = getattr(value, argument) + field = getattr(value, argument) # pyrefly: ignore[bad-argument-type] if not isinstance(field, (config_lib.Buildable, list, tuple, dict)): field_repr = repr(field) if len(field_repr) > threshold: @@ -332,7 +332,7 @@ def traverse(value, state: daglish.State): repr(field), width=threshold, placeholder='...' ) prefix = _TruncatedRepr(s) - setattr(value, argument, prefix) + setattr(value, argument, prefix) # pyrefly: ignore[bad-argument-type] return state.map_children(value) return daglish.MemoizedTraversal.run(traverse, config) diff --git a/fiddle/_src/experimental/yaml_serialization.py b/fiddle/_src/experimental/yaml_serialization.py index 4a12e57d..388faa26 100644 --- a/fiddle/_src/experimental/yaml_serialization.py +++ b/fiddle/_src/experimental/yaml_serialization.py @@ -54,7 +54,7 @@ def _config_representer(dumper, data, type_name="fdl.Config"): ) value["__fn_or_cls__"] = { - "module": inspect.getmodule(config_lib.get_callable(data)).__name__, + "module": inspect.getmodule(config_lib.get_callable(data)).__name__, # pyrefly: ignore[missing-attribute] "name": config_lib.get_callable(data).__qualname__, } return dumper.represent_mapping(f"!{type_name}", value) diff --git a/fiddle/_src/extensions/tf.py b/fiddle/_src/extensions/tf.py index 75900009..dac9f53a 100644 --- a/fiddle/_src/extensions/tf.py +++ b/fiddle/_src/extensions/tf.py @@ -65,9 +65,9 @@ def make_dtype_name(module, dtype_name): @py_val_to_cst_converter.register_py_val_to_cst_converter(is_tensor) def convert_tensor_to_cst(value, convert_child): return cst.Call( - func=cst.Attribute(value=convert_child(tf), attr=cst.Name("constant")), + func=cst.Attribute(value=convert_child(tf), attr=cst.Name("constant")), # pyrefly: ignore[bad-argument-type] args=[ - cst.Arg(convert_child(value.numpy().tolist())), + cst.Arg(convert_child(value.numpy().tolist())), # pyrefly: ignore[bad-argument-type] py_val_to_cst_converter.kwarg_to_cst("dtype", convert_child(value.dtype)), py_val_to_cst_converter.kwarg_to_cst( @@ -76,12 +76,12 @@ def convert_tensor_to_cst(value, convert_child): @py_val_to_cst_converter.register_py_val_to_cst_converter(tf.DType) def convert_dtype_to_cst(value, convert_child): - return cst.Attribute(value=convert_child(tf), attr=cst.Name(value.name)) + return cst.Attribute(value=convert_child(tf), attr=cst.Name(value.name)) # pyrefly: ignore[bad-argument-type] @py_val_to_cst_converter.register_py_val_to_cst_converter(tf.TensorShape) def convert_tensor_shape_to_cst(value, convert_child): shape_list = None if value.rank is None else value.as_list() return cst.Call( func=cst.Attribute( - value=convert_child(tf), attr=cst.Name("TensorShape")), - args=[cst.Arg(convert_child(shape_list))]) + value=convert_child(tf), attr=cst.Name("TensorShape")), # pyrefly: ignore[bad-argument-type] + args=[cst.Arg(convert_child(shape_list))]) # pyrefly: ignore[bad-argument-type] diff --git a/fiddle/_src/extensions/tf_test.py b/fiddle/_src/extensions/tf_test.py index 16e77c56..9d87399d 100644 --- a/fiddle/_src/extensions/tf_test.py +++ b/fiddle/_src/extensions/tf_test.py @@ -129,7 +129,7 @@ def test_default_printing(self): ]) def test_py_val_to_cst_converter(self, value, expected): cst_expr = py_val_to_cst_converter.convert_py_val_to_cst(value) - cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) + cst_module = cst.Module([cst.SimpleStatementLine([cst.Expr(cst_expr)])]) # pyrefly: ignore[bad-argument-type] self.assertEqual(cst_module.code.strip(), expected) def test_serialization(self): diff --git a/fiddle/_src/testing/autotest.py b/fiddle/_src/testing/autotest.py index 0443e11b..3d44d1a9 100644 --- a/fiddle/_src/testing/autotest.py +++ b/fiddle/_src/testing/autotest.py @@ -68,8 +68,8 @@ def test_base_config(self, module, fn_name): _ = building.build(cfg) fn = functools.partialmethod(test_base_config, module, name) - fn.__name__ = f'test_{name}' - test_functions[fn.__name__] = fn + fn.__name__ = f'test_{name}' # pyrefly: ignore[missing-attribute] + test_functions[fn.__name__] = fn # pyrefly: ignore[missing-attribute] # Test all combination of fiddlers and base configurations by default. for base_name in module_reflection.find_base_config_like_things(module): @@ -87,8 +87,8 @@ def test_base_and_fiddler(self, module, base_name, fiddler_name): fn = functools.partialmethod(test_base_and_fiddler, module, base_name, fiddler_name) - fn.__name__ = f'test_{base_name}_and_{fiddler_name}' - test_functions[fn.__name__] = fn + fn.__name__ = f'test_{base_name}_and_{fiddler_name}' # pyrefly: ignore[missing-attribute] + test_functions[fn.__name__] = fn # pyrefly: ignore[missing-attribute] if (not module_reflection.find_base_config_like_things(module) and module_reflection.find_fiddler_like_things(module)): @@ -141,7 +141,7 @@ def load_tests(loader: unittest.TestLoader, tests: unittest.TestSuite, """ del tests # Unused. del pattern # Unused. - module = load_module_from_path(_FLAG_FIDDLE_CONFIG_MODULE.value) + module = load_module_from_path(_FLAG_FIDDLE_CONFIG_MODULE.value) # pyrefly: ignore[bad-argument-type] suite = load_tests_from_module( loader, module, skip_building=_FLAG_SKIP_BUILDING.value) return suite diff --git a/fiddle/_src/testing/example/fake_encoder_decoder.py b/fiddle/_src/testing/example/fake_encoder_decoder.py index 43f846d8..966b818f 100644 --- a/fiddle/_src/testing/example/fake_encoder_decoder.py +++ b/fiddle/_src/testing/example/fake_encoder_decoder.py @@ -84,7 +84,7 @@ def fixture(kernel_init: str = "uniform()"): shared_token_embedder = TokenEmbedder(dtype) return FakeEncoderDecoder( encoder=FakeEncoder( - embedders={ + embedders={ # pyrefly: ignore[bad-argument-type] "tokens": shared_token_embedder, "position": None }, diff --git a/fiddle/_src/testing/nested_values.py b/fiddle/_src/testing/nested_values.py index 1427a84e..5822fbb8 100644 --- a/fiddle/_src/testing/nested_values.py +++ b/fiddle/_src/testing/nested_values.py @@ -138,7 +138,7 @@ def generate_buildable(): return buildable def generate_alias(): - for value in enumerate(share_objects): + for value in enumerate(share_objects): # pyrefly: ignore[bad-argument-type] if calculate_nested_value_depth(value) < max_depth: return value return generate_value() diff --git a/fiddle/_src/validation/check_types.py b/fiddle/_src/validation/check_types.py index dafa8977..aeb5c3de 100644 --- a/fiddle/_src/validation/check_types.py +++ b/fiddle/_src/validation/check_types.py @@ -57,7 +57,7 @@ def _check_types_recursive(value: Any, state: daglish.State) -> Any: # generics raise a `TypeError`. For these, get their origin type and # check that the origin types match. type_hint_origin = get_origin(type_hints[arg_name]) - if not isinstance(arg_value, type_hint_origin): + if not isinstance(arg_value, type_hint_origin): # pyrefly: ignore[bad-argument-type] add_error = True if add_error: path_str = daglish.path_str(state.current_path) diff --git a/fiddle/_src/validation/check_types_test.py b/fiddle/_src/validation/check_types_test.py index 54d60eb3..8b566fc4 100644 --- a/fiddle/_src/validation/check_types_test.py +++ b/fiddle/_src/validation/check_types_test.py @@ -69,7 +69,7 @@ class StackedEncoder: encoders: List[fake_encoder_decoder.FakeEncoder] -class CheckTypesTest(absltest.TestCase, unittest.TestCase): +class CheckTypesTest(absltest.TestCase, unittest.TestCase): # pyrefly: ignore[inconsistent-inheritance] def test_type_validation(self): cfg = config.Config(Experiment, dataset=BadDataset()) diff --git a/fiddle/_src/validation/no_custom_objects.py b/fiddle/_src/validation/no_custom_objects.py index 3aec72c5..9a6e9590 100644 --- a/fiddle/_src/validation/no_custom_objects.py +++ b/fiddle/_src/validation/no_custom_objects.py @@ -74,7 +74,7 @@ def get_config_errors(config: Any) -> List[str]: errors = [] def history_str(state): - return ", " + _concise_history(_get_history_from_state(state)) + return ", " + _concise_history(_get_history_from_state(state)) # pyrefly: ignore[bad-argument-type] def traverse(value, state: daglish.State): path_str = daglish.path_str(state.current_path)