diff --git a/jax3d/projects/generative/nerf/autoencoder/data.py b/jax3d/projects/generative/nerf/autoencoder/data.py index c911cae..6f0d71c 100644 --- a/jax3d/projects/generative/nerf/autoencoder/data.py +++ b/jax3d/projects/generative/nerf/autoencoder/data.py @@ -26,9 +26,9 @@ @dataclasses.dataclass(frozen=True) class TFDSImageDatasetReader(): """Dataset reader wrapping TFDS image datasets.""" - dataset_name: str = gin.REQUIRED - resolution: int = gin.REQUIRED - batch_size: int = gin.REQUIRED + dataset_name: str = gin.REQUIRED # pyrefly: ignore[bad-assignment] + resolution: int = gin.REQUIRED # pyrefly: ignore[bad-assignment] + batch_size: int = gin.REQUIRED # pyrefly: ignore[bad-assignment] split: str = "train" eval_fraction: float = 0.05 @@ -82,9 +82,9 @@ def extract_image(data): @dataclasses.dataclass(frozen=True) class TiledMNISTDatasetReader(): """Dataset reader wrapping TFDS image datasets.""" - resolution: int = gin.REQUIRED - batch_size: int = gin.REQUIRED - tile_factor: int = gin.REQUIRED + resolution: int = gin.REQUIRED # pyrefly: ignore[bad-assignment] + batch_size: int = gin.REQUIRED # pyrefly: ignore[bad-assignment] + tile_factor: int = gin.REQUIRED # pyrefly: ignore[bad-assignment] split: str = "train" eval_fraction: float = 0.05 diff --git a/jax3d/projects/generative/nerf/autoencoder/loss.py b/jax3d/projects/generative/nerf/autoencoder/loss.py index d741475..4b61f6b 100644 --- a/jax3d/projects/generative/nerf/autoencoder/loss.py +++ b/jax3d/projects/generative/nerf/autoencoder/loss.py @@ -30,10 +30,10 @@ def transformer_loss_fn( model_parameters: models.ModelParameters, data: Dict[str, Any], - rng: PRNGKey, + rng: PRNGKey, # pyrefly: ignore[not-a-type] step: jnp.ndarray, pixel_batch_size=512, -) -> Tuple[FloatArray, Dict[str, FloatArray]]: +) -> Tuple[FloatArray, Dict[str, FloatArray]]: # pyrefly: ignore[not-a-type] """The main autoencoder loss function. Args: @@ -92,7 +92,7 @@ def take(pixels, inds): total_loss += reconstruction_loss * reconstruction_loss_weight # Always compute PSNR on gamma values for consistency - psnr = metrics.psnr(predicted_rgb, gt_rgb) + psnr = metrics.psnr(predicted_rgb, gt_rgb) # pyrefly: ignore[bad-argument-type] loss_terms["Training PSNR"] = psnr return total_loss, loss_terms diff --git a/jax3d/projects/generative/nerf/autoencoder/models.py b/jax3d/projects/generative/nerf/autoencoder/models.py index ff4bdd9..e030c43 100644 --- a/jax3d/projects/generative/nerf/autoencoder/models.py +++ b/jax3d/projects/generative/nerf/autoencoder/models.py @@ -246,7 +246,7 @@ def __call__(self, image, pixels=None, step=None, is_training=False): return rgb - def initialize_parameters(self, rng_key: PRNGKey, + def initialize_parameters(self, rng_key: PRNGKey, # pyrefly: ignore[not-a-type] image_size) -> ModelParameters: batch_size = 7 pixel_batch_size = 11 diff --git a/jax3d/projects/generative/nerf/autoencoder/trainer.py b/jax3d/projects/generative/nerf/autoencoder/trainer.py index ecb45e2..935582a 100644 --- a/jax3d/projects/generative/nerf/autoencoder/trainer.py +++ b/jax3d/projects/generative/nerf/autoencoder/trainer.py @@ -39,11 +39,11 @@ class TransformerTrainState(trainer.TrainState): """Training state for the NeRF Transformer Model.""" model_parameters: models.ModelParameters optimizer_state: optax.OptState - rng: PRNGKey + rng: PRNGKey # pyrefly: ignore[not-a-type] def to_serializable(self) -> "TransformerTrainState": """Transforms the state values into a form suitable for serialization.""" - return self.replace( + return self.replace( # pyrefly: ignore[missing-attribute] model_parameters=jax_utils.unreplicate(self.model_parameters), optimizer_state=jax_utils.unreplicate(self.optimizer_state), ) @@ -51,7 +51,7 @@ def to_serializable(self) -> "TransformerTrainState": def from_serializable(self) -> "TransformerTrainState": """Transforms deserialized values into a form suitable for training.""" state = self - state = state.replace( + state = state.replace( # pyrefly: ignore[missing-attribute] model_parameters=jax_utils.replicate(state.model_parameters), optimizer_state=jax_utils.replicate(state.optimizer_state), ) @@ -89,7 +89,7 @@ def data_reader(self) -> "dataset_reader_class": def init_state(self) -> TransformerTrainState: """Initializes training state.""" - rng = jax.random.PRNGKey(self.random_seed) + rng = jax.random.PRNGKey(self.random_seed) # pyrefly: ignore[bad-argument-type] train_rng, model_init_rng = jax.random.split(rng) res = self.data_reader.resolution model_parameters = self.model.initialize_parameters(model_init_rng, @@ -135,7 +135,7 @@ def _loss_fn(model_parameters, data, rng, step): return jax.pmap(per_device_train_step, "replicas") - def train_step( + def train_step( # pyrefly: ignore[bad-override] self, train_state: TransformerTrainState, inputs: Dict[str, Any], scratch: Optional[Dict[str, Any]] ) -> Tuple[TransformerTrainState, Dict[str, Any], Dict[str, Any]]: @@ -173,7 +173,7 @@ def train_step( loss_terms = jax.tree.map(np.array, loss_terms) loss_terms["param_count"] = scratch["param_count"] - new_train_state = train_state.replace( + new_train_state = train_state.replace( # pyrefly: ignore[missing-attribute] model_parameters=new_model_parameters, optimizer_state=new_optimizer_state, rng=next_rng, @@ -181,7 +181,7 @@ def train_step( return new_train_state, loss_terms, scratch - def eval_step(self, train_state: TransformerTrainState, + def eval_step(self, train_state: TransformerTrainState, # pyrefly: ignore[bad-override] summary_writer: tensorboard.SummaryWriter, scratch: Optional[Any]) -> Any: """Performs evaluation on a model checkpoint. diff --git a/jax3d/projects/generative/nerf/camera.py b/jax3d/projects/generative/nerf/camera.py index d000a6d..c40dee4 100644 --- a/jax3d/projects/generative/nerf/camera.py +++ b/jax3d/projects/generative/nerf/camera.py @@ -150,8 +150,8 @@ def make_camera( # pytype: disable=annotation-type-mismatch # jax-ndarray focal_length: Union[jnp.ndarray, float], principal_point: jnp.ndarray, image_size: jnp.ndarray, - skew: Union[jnp.ndarray, float] = None, - pixel_aspect_ratio: Union[jnp.ndarray, float] = None, + skew: Union[jnp.ndarray, float] = None, # pyrefly: ignore[bad-function-definition] + pixel_aspect_ratio: Union[jnp.ndarray, float] = None, # pyrefly: ignore[bad-function-definition] radial_distortion: Optional[jnp.ndarray] = None, tangential_distortion: Optional[jnp.ndarray] = None) -> CameraType: """Create a dictionary containing standard values for representing a camera. @@ -222,7 +222,7 @@ def look_at(eye_position: jnp.ndarray, # +Z = forward camera_forward = target - eye_position camera_forward /= jnp.linalg.norm(camera_forward) - camera_right = jnp.cross(camera_forward, global_up) + camera_right = jnp.cross(camera_forward, global_up) # pyrefly: ignore[bad-argument-type] camera_right /= jnp.linalg.norm(camera_right) camera_up = jnp.cross(camera_right, camera_forward) orientation = jnp.stack([-camera_right, camera_up, camera_forward], axis=0) diff --git a/jax3d/projects/generative/nerf/glo_nerf/eval.py b/jax3d/projects/generative/nerf/glo_nerf/eval.py index 0fb2c6c..66697b0 100644 --- a/jax3d/projects/generative/nerf/glo_nerf/eval.py +++ b/jax3d/projects/generative/nerf/glo_nerf/eval.py @@ -179,7 +179,7 @@ def flatten_views(t): render = models.Model().apply( model_parameters, inputs, rays, rng=rng, step=step) - pred = render["gamma_rgb"] + pred = render["gamma_rgb"] # pyrefly: ignore[bad-index] gt = data_flat["gamma_rgb"] if apply_mask: pred *= data_flat["weight"] diff --git a/jax3d/projects/generative/nerf/glo_nerf/loss.py b/jax3d/projects/generative/nerf/glo_nerf/loss.py index da16460..92c8fe1 100644 --- a/jax3d/projects/generative/nerf/glo_nerf/loss.py +++ b/jax3d/projects/generative/nerf/glo_nerf/loss.py @@ -32,11 +32,11 @@ def transformer_nerf_loss_fn( model_parameters: models.ModelParameters, inputs: models.ModelInputs, data: Dict[str, Any], - rng: PRNGKey, + rng: PRNGKey, # pyrefly: ignore[not-a-type] step: jnp.ndarray, mask_mode: str = "alpha_supervision", reconstruct_gamma_rgb: bool = True -) -> Tuple[FloatArray, Dict[str, FloatArray]]: +) -> Tuple[FloatArray, Dict[str, FloatArray]]: # pyrefly: ignore[not-a-type] """The main GLO NeRF loss function including all terms. Args: @@ -65,7 +65,7 @@ def flatten_views(t): data = jax.tree.map(flatten_views, data) latent_tokens = inputs.latent_tokens latent_tokens = jax.tree.map(flatten_views, latent_tokens) - inputs = inputs.replace(latent_tokens=latent_tokens) + inputs = inputs.replace(latent_tokens=latent_tokens) # pyrefly: ignore[missing-attribute] origins, directions = jax.vmap(jax_camera.pixels_to_rays)( data["camera"], data["pixel_coordinates"]) @@ -96,10 +96,10 @@ def flatten_views(t): # gamma decoding their images this is effectively the loss they are using. # (https://en.wikipedia.org/wiki/Gamma_correction). gt_rgb = data["gamma_rgb"] - predicted_rgb = render["gamma_rgb"] + predicted_rgb = render["gamma_rgb"] # pyrefly: ignore[bad-index] else: gt_rgb = image_utility.srgb_gamma_to_linear(data["gamma_rgb"]) - predicted_rgb = render["linear_rgb"] + predicted_rgb = render["linear_rgb"] # pyrefly: ignore[bad-index] if mask_mode == "multiply": gt_rgb *= foreground_mask @@ -115,19 +115,19 @@ def flatten_views(t): gt_rgb = data["gamma_rgb"] if mask_mode == "multiply": gt_rgb *= foreground_mask - psnr = metrics.psnr(render["gamma_rgb"], gt_rgb) + psnr = metrics.psnr(render["gamma_rgb"], gt_rgb) # pyrefly: ignore[bad-index] loss_terms["Training PSNR"] = psnr if mask_mode == "alpha_supervision": with gin.config_scope("alpha"): alpha_loss, alpha_loss_weight = losses.reconstruction( - foreground_mask, render["alpha"]) + foreground_mask, render["alpha"]) # pyrefly: ignore[bad-index] if alpha_loss_weight != 0.0: loss_terms["Alpha"] = alpha_loss total_loss += alpha_loss_weight * alpha_loss hard_surface_loss, hard_surface_loss_weight = losses.hard_surface( - render["sample_weights"]) + render["sample_weights"]) # pyrefly: ignore[bad-index] loss_terms["Hard Surface"] = hard_surface_loss total_loss += hard_surface_loss_weight * hard_surface_loss diff --git a/jax3d/projects/generative/nerf/glo_nerf/models.py b/jax3d/projects/generative/nerf/glo_nerf/models.py index baf2966..70ff272 100644 --- a/jax3d/projects/generative/nerf/glo_nerf/models.py +++ b/jax3d/projects/generative/nerf/glo_nerf/models.py @@ -313,11 +313,11 @@ def setup(self): def evaluate_nerf( self, - points: FloatArray["N", ..., 3], - directions: FloatArray["N", ..., 3], + points: FloatArray["N", ..., 3], # pyrefly: ignore[not-a-type, unknown-name] + directions: FloatArray["N", ..., 3], # pyrefly: ignore[not-a-type, unknown-name] inputs: ModelInputs, step: int, - noise_rng: Optional[PRNGKey] = None) -> Dict[str, FloatArray]: + noise_rng: Optional[PRNGKey] = None) -> Dict[str, FloatArray]: # pyrefly: ignore[not-a-type] """NeRF decoder forward pass.""" density, rgb = self.decoder(points, inputs.latent_tokens) @@ -325,8 +325,8 @@ def evaluate_nerf( return results def surface_normal_from_density( - self, points: FloatArray["N", ..., 3], inputs, step: int - ) -> FloatArray["N", ..., 3]: + self, points: FloatArray["N", ..., 3], inputs, step: int # pyrefly: ignore[not-a-type, unknown-name] + ) -> FloatArray["N", ..., 3]: # pyrefly: ignore[not-a-type, unknown-name] """Compute the normals of the density field at the given points.""" # Input points are of the form [ N, ..., 3]. Compute the product of the @@ -349,7 +349,7 @@ def density_fn(fn_points): def __call__(self, inputs: ModelInputs, - rays: FloatArray["K R", 6], + rays: FloatArray["K R", 6], # pyrefly: ignore[not-a-type] rng=None, near=None, far=None, @@ -404,7 +404,7 @@ def __call__(self, initial_nerf_result = self.evaluate_nerf(initial_sample_coordinates, initial_sample_directions, inputs, - step, initial_noise_rng) + step, initial_noise_rng) # pyrefly: ignore[bad-argument-type] initial_render_results = volume_rendering.volume_rendering( sample_values=(), # We ignore the accumulated RGB from initial samples @@ -448,7 +448,7 @@ def __call__(self, combined_nerf_result = self.evaluate_nerf(combined_sample_coordinates, combined_sample_directions, - inputs, step, + inputs, step, # pyrefly: ignore[bad-argument-type] importance_noise_rng) else: @@ -473,7 +473,7 @@ def __call__(self, importance_nerf_result = self.evaluate_nerf(importance_sample_coordinates, importance_sample_directions, - inputs, step, + inputs, step, # pyrefly: ignore[bad-argument-type] importance_noise_rng) # We first append importance sample depths and values along the sample @@ -556,7 +556,7 @@ def __call__(self, surface_points = origins + directions * expected_depth return_values["depth"] = expected_depth - normals = self.surface_normal_from_density(surface_points, inputs, step) + normals = self.surface_normal_from_density(surface_points, inputs, step) # pyrefly: ignore[bad-argument-type] return_values["analytic_normal"] = normals if return_additional_sample_data: @@ -565,7 +565,7 @@ def __call__(self, alpha = render_results.ray_alpha[..., None] return_values["alpha"] = alpha - foreground_rgb = render_results.ray_values["rgb"] + foreground_rgb = render_results.ray_values["rgb"] # pyrefly: ignore[bad-index] if self.use_background_model: background_latent = inputs.latent_tokens[..., 0, :] @@ -574,7 +574,7 @@ def __call__(self, # need to add it to the attenuated background value. pixel_rgb = foreground_rgb + (1.0 - alpha) * background_rgb else: - background_rgb = jnp.zeros_like(foreground_rgb) + background_rgb = jnp.zeros_like(foreground_rgb) # pyrefly: ignore[bad-argument-type] pixel_rgb = foreground_rgb # Up until this point RGB values could be interpreted a gamma encoded or @@ -726,13 +726,13 @@ def unpad(tensor): for key in results: result_i = jnp.concatenate(results[key]) - results[key] = result_i.reshape(height, width, *result_i.shape[1:]) + results[key] = result_i.reshape(height, width, *result_i.shape[1:]) # pyrefly: ignore[unsupported-operation] return results return render_image - def initialize_parameters(self, rng_key: PRNGKey, num_tokens, + def initialize_parameters(self, rng_key: PRNGKey, num_tokens, # pyrefly: ignore[not-a-type] token_dim) -> ModelParameters: batch_size = 7 num_rays = 13 # Prime numbers to help catch shape errors. diff --git a/jax3d/projects/generative/nerf/glo_nerf/trainer.py b/jax3d/projects/generative/nerf/glo_nerf/trainer.py index a788df0..86b4c96 100644 --- a/jax3d/projects/generative/nerf/glo_nerf/trainer.py +++ b/jax3d/projects/generative/nerf/glo_nerf/trainer.py @@ -40,11 +40,11 @@ class TransformerNeRFTrainState(trainer.TrainState): model_parameters: models.ModelParameters latent_table: np.ndarray optimizer_state: optax.OptState - rng: PRNGKey + rng: PRNGKey # pyrefly: ignore[not-a-type] def to_serializable(self) -> "TransformerNeRFTrainState": """Transforms the state values into a form suitable for serialization.""" - return self.replace( + return self.replace( # pyrefly: ignore[missing-attribute] model_parameters=jax_utils.unreplicate(self.model_parameters), optimizer_state=jax_utils.unreplicate(self.optimizer_state), ) @@ -52,7 +52,7 @@ def to_serializable(self) -> "TransformerNeRFTrainState": def from_serializable(self) -> "TransformerNeRFTrainState": """Transforms deserialized values into a form suitable for training.""" state = self - state = state.replace( + state = state.replace( # pyrefly: ignore[missing-attribute] latent_table=np.copy(state.latent_table), model_parameters=jax_utils.replicate(state.model_parameters), optimizer_state=jax_utils.replicate(state.optimizer_state), @@ -99,7 +99,7 @@ def init_state(self) -> TransformerNeRFTrainState: latent_table = np.zeros((self.data_reader.identity_count, self.num_latent_tokens, self.latent_token_dim)) - rng = jax.random.PRNGKey(self.random_seed) + rng = jax.random.PRNGKey(self.random_seed) # pyrefly: ignore[bad-argument-type] train_rng, model_init_rng = jax.random.split(rng) model_parameters = self.model.initialize_parameters(model_init_rng, self.num_latent_tokens, @@ -155,7 +155,7 @@ def _loss_fn(params, data, rng, step): return jax.pmap(per_device_train_step, "replicas") - def train_step( + def train_step( # pyrefly: ignore[bad-override] self, train_state: TransformerNeRFTrainState, inputs: Dict[str, Any], scratch: Optional[Dict[str, Any]] ) -> Tuple[TransformerNeRFTrainState, Dict[str, Any], Dict[str, Any]]: @@ -199,7 +199,7 @@ def train_step( steps = -latent_learning_rate * latent_grad train_state.latent_table[latent_ids] += steps - new_train_state = train_state.replace( + new_train_state = train_state.replace( # pyrefly: ignore[missing-attribute] model_parameters=new_model_parameters, optimizer_state=new_optimizer_state, latent_table=train_state.latent_table, @@ -208,7 +208,7 @@ def train_step( return new_train_state, loss_terms, scratch - def eval_step(self, train_state: TransformerNeRFTrainState, + def eval_step(self, train_state: TransformerNeRFTrainState, # pyrefly: ignore[bad-override] summary_writer: tensorboard.SummaryWriter, scratch: Optional[Any]) -> Any: """Performs evaluation on a model checkpoint. diff --git a/jax3d/projects/generative/nerf/lightfield/eval.py b/jax3d/projects/generative/nerf/lightfield/eval.py index 0b6a7f9..afe9804 100644 --- a/jax3d/projects/generative/nerf/lightfield/eval.py +++ b/jax3d/projects/generative/nerf/lightfield/eval.py @@ -179,7 +179,7 @@ def flatten_views(t): render = models.Model().apply( model_parameters, inputs, rays, rng=rng, step=step) - pred = render["gamma_rgb"] + pred = render["gamma_rgb"] # pyrefly: ignore[bad-index] gt = data_flat["gamma_rgb"] if apply_mask: pred *= data_flat["weight"] diff --git a/jax3d/projects/generative/nerf/lightfield/loss.py b/jax3d/projects/generative/nerf/lightfield/loss.py index 38b272c..07f643c 100644 --- a/jax3d/projects/generative/nerf/lightfield/loss.py +++ b/jax3d/projects/generative/nerf/lightfield/loss.py @@ -33,11 +33,11 @@ def lightfield_loss_fn( model_parameters: models.ModelParameters, inputs: models.ModelInputs, data: Dict[str, Any], - rng: PRNGKey, + rng: PRNGKey, # pyrefly: ignore[not-a-type] step: jnp.ndarray, mask_mode: str = "none", reconstruct_gamma_rgb: bool = True -) -> Tuple[FloatArray, Dict[str, FloatArray]]: +) -> Tuple[FloatArray, Dict[str, FloatArray]]: # pyrefly: ignore[not-a-type] """The main GLO NeRF loss function including all terms. Args: @@ -66,7 +66,7 @@ def flatten_views(t): data = jax.tree.map(flatten_views, data) latent_tokens = inputs.latent_tokens latent_tokens = jax.tree.map(flatten_views, latent_tokens) - inputs = inputs.replace(latent_tokens=latent_tokens) + inputs = inputs.replace(latent_tokens=latent_tokens) # pyrefly: ignore[missing-attribute] origins, directions = jax.vmap(jax_camera.pixels_to_rays)( data["camera"], data["pixel_coordinates"]) @@ -96,10 +96,10 @@ def flatten_views(t): # gamma decoding their images this is effectively the loss they are using. # (https://en.wikipedia.org/wiki/Gamma_correction). gt_rgb = data["gamma_rgb"] - predicted_rgb = render["gamma_rgb"] + predicted_rgb = render["gamma_rgb"] # pyrefly: ignore[bad-index] else: gt_rgb = image_utility.srgb_gamma_to_linear(data["gamma_rgb"]) - predicted_rgb = render["linear_rgb"] + predicted_rgb = render["linear_rgb"] # pyrefly: ignore[bad-index] if mask_mode == "multiply": gt_rgb *= foreground_mask @@ -115,7 +115,7 @@ def flatten_views(t): gt_rgb = data["gamma_rgb"] if mask_mode == "multiply": gt_rgb *= foreground_mask - psnr = metrics.psnr(render["gamma_rgb"], gt_rgb) + psnr = metrics.psnr(render["gamma_rgb"], gt_rgb) # pyrefly: ignore[bad-index] loss_terms["Training PSNR"] = psnr return total_loss, loss_terms diff --git a/jax3d/projects/generative/nerf/lightfield/models.py b/jax3d/projects/generative/nerf/lightfield/models.py index a030050..801987d 100644 --- a/jax3d/projects/generative/nerf/lightfield/models.py +++ b/jax3d/projects/generative/nerf/lightfield/models.py @@ -274,7 +274,7 @@ def setup(self): def __call__(self, inputs: ModelInputs, - rays: FloatArray["K R", 6], + rays: FloatArray["K R", 6], # pyrefly: ignore[not-a-type] rng=None, step=None, is_training=False): @@ -387,13 +387,13 @@ def unpad(tensor): for key in results: result_i = jnp.concatenate(results[key]) - results[key] = result_i.reshape(height, width, *result_i.shape[1:]) + results[key] = result_i.reshape(height, width, *result_i.shape[1:]) # pyrefly: ignore[unsupported-operation] return results return render_image - def initialize_parameters(self, rng_key: PRNGKey, num_tokens, + def initialize_parameters(self, rng_key: PRNGKey, num_tokens, # pyrefly: ignore[not-a-type] token_dim) -> ModelParameters: batch_size = 7 num_rays = 13 # Prime numbers to help catch shape errors. diff --git a/jax3d/projects/generative/nerf/lightfield/trainer.py b/jax3d/projects/generative/nerf/lightfield/trainer.py index dc36210..5e267ba 100644 --- a/jax3d/projects/generative/nerf/lightfield/trainer.py +++ b/jax3d/projects/generative/nerf/lightfield/trainer.py @@ -40,11 +40,11 @@ class TransformerNeRFTrainState(trainer.TrainState): model_parameters: models.ModelParameters latent_table: np.ndarray optimizer_state: optax.OptState - rng: PRNGKey + rng: PRNGKey # pyrefly: ignore[not-a-type] def to_serializable(self) -> "TransformerNeRFTrainState": """Transforms the state values into a form suitable for serialization.""" - return self.replace( + return self.replace( # pyrefly: ignore[missing-attribute] model_parameters=jax_utils.unreplicate(self.model_parameters), optimizer_state=jax_utils.unreplicate(self.optimizer_state), ) @@ -52,7 +52,7 @@ def to_serializable(self) -> "TransformerNeRFTrainState": def from_serializable(self) -> "TransformerNeRFTrainState": """Transforms deserialized values into a form suitable for training.""" state = self - state = state.replace( + state = state.replace( # pyrefly: ignore[missing-attribute] latent_table=np.copy(state.latent_table), model_parameters=jax_utils.replicate(state.model_parameters), optimizer_state=jax_utils.replicate(state.optimizer_state), @@ -99,7 +99,7 @@ def init_state(self) -> TransformerNeRFTrainState: latent_table = np.zeros((self.data_reader.identity_count, self.num_latent_tokens, self.latent_token_dim)) - rng = jax.random.PRNGKey(self.random_seed) + rng = jax.random.PRNGKey(self.random_seed) # pyrefly: ignore[bad-argument-type] train_rng, model_init_rng = jax.random.split(rng) model_parameters = self.model.initialize_parameters(model_init_rng, self.num_latent_tokens, @@ -155,7 +155,7 @@ def _loss_fn(params, data, rng, step): return jax.pmap(per_device_train_step, "replicas") - def train_step( + def train_step( # pyrefly: ignore[bad-override] self, train_state: TransformerNeRFTrainState, inputs: Dict[str, Any], scratch: Optional[Dict[str, Any]] ) -> Tuple[TransformerNeRFTrainState, Dict[str, Any], Dict[str, Any]]: @@ -209,7 +209,7 @@ def train_step( steps = -latent_learning_rate * latent_grad train_state.latent_table[latent_ids] += steps - new_train_state = train_state.replace( + new_train_state = train_state.replace( # pyrefly: ignore[missing-attribute] model_parameters=new_model_parameters, optimizer_state=new_optimizer_state, latent_table=train_state.latent_table, @@ -218,7 +218,7 @@ def train_step( return new_train_state, loss_terms, scratch - def eval_step(self, train_state: TransformerNeRFTrainState, + def eval_step(self, train_state: TransformerNeRFTrainState, # pyrefly: ignore[bad-override] summary_writer: tensorboard.SummaryWriter, scratch: Optional[Any]) -> Any: """Performs evaluation on a model checkpoint. diff --git a/jax3d/projects/generative/nerf/losses.py b/jax3d/projects/generative/nerf/losses.py index b066a5f..440ae75 100644 --- a/jax3d/projects/generative/nerf/losses.py +++ b/jax3d/projects/generative/nerf/losses.py @@ -38,10 +38,10 @@ def _get_norm_fn(name: str) -> Callable[[jnp.ndarray], jnp.ndarray]: @gin.configurable(allowlist=["low_threshold", "high_threshold"]) -def tri_mode_clipping(ground_truth: FloatArray[..., "C"], - predicted: FloatArray[..., "C"], +def tri_mode_clipping(ground_truth: FloatArray[..., "C"], # pyrefly: ignore[not-a-type, unknown-name] + predicted: FloatArray[..., "C"], # pyrefly: ignore[not-a-type, unknown-name] low_threshold: float = 0.0, - high_threshold: float = 1.0) -> FloatArray[..., "C"]: + high_threshold: float = 1.0) -> FloatArray[..., "C"]: # pyrefly: ignore[not-a-type, unknown-name] """An error clipping scheme for data saturated outside a given range. For ground truth pixels outside this range, predicted pixels only affect the @@ -81,12 +81,12 @@ def tri_mode_clipping(ground_truth: FloatArray[..., "C"], @gin.configurable( "reconstruction_loss", allowlist=["weight", "norm", "use_tri_mode"]) -def reconstruction(ground_truth: FloatArray[..., "C"], - predicted: FloatArray[..., "C"], - mask: Optional[FloatArray[...]] = None, +def reconstruction(ground_truth: FloatArray[..., "C"], # pyrefly: ignore[not-a-type, unknown-name] + predicted: FloatArray[..., "C"], # pyrefly: ignore[not-a-type, unknown-name] + mask: Optional[FloatArray[...]] = None, # pyrefly: ignore[not-a-type] weight: float = 1.0, norm: str = "l2", - use_tri_mode: bool = False) -> Tuple[FloatArray, float]: + use_tri_mode: bool = False) -> Tuple[FloatArray, float]: # pyrefly: ignore[not-a-type] """A photometric reconstruction loss. Args: @@ -118,12 +118,12 @@ def reconstruction(ground_truth: FloatArray[..., "C"], @gin.configurable( "normal_consistency_loss", allowlist=["weight", "mode", "hold_analytic_normals_constant"]) -def normal_consistency(analytic_normals: FloatArray[..., 3], - predicted_normals: FloatArray[..., 3], - mask: Optional[FloatArray[...]] = None, +def normal_consistency(analytic_normals: FloatArray[..., 3], # pyrefly: ignore[not-a-type] + predicted_normals: FloatArray[..., 3], # pyrefly: ignore[not-a-type] + mask: Optional[FloatArray[...]] = None, # pyrefly: ignore[not-a-type] weight: float = 0.0, hold_analytic_normals_constant: bool = True, - mode: str = "error") -> Tuple[FloatArray, float]: + mode: str = "error") -> Tuple[FloatArray, float]: # pyrefly: ignore[not-a-type] """Loss for enforcing consistency between predicted and analytic normals. Args: @@ -158,9 +158,9 @@ def normal_consistency(analytic_normals: FloatArray[..., 3], @gin.configurable("color_correction_regularization", allowlist=["weight"]) -def color_correction_regularization(error: FloatArray[...], +def color_correction_regularization(error: FloatArray[...], # pyrefly: ignore[not-a-type] weight: float = 0.0 - ) -> Tuple[FloatArray, float]: + ) -> Tuple[FloatArray, float]: # pyrefly: ignore[not-a-type] """Color correction regularization. Args: @@ -175,8 +175,8 @@ def color_correction_regularization(error: FloatArray[...], @gin.configurable("hard_surface_loss", allowlist=["weight"]) -def hard_surface(sample_weights: FloatArray[...], - weight: float = 0.0) -> Tuple[FloatArray, float]: +def hard_surface(sample_weights: FloatArray[...], # pyrefly: ignore[not-a-type] + weight: float = 0.0) -> Tuple[FloatArray, float]: # pyrefly: ignore[not-a-type] """Hard surface density regularizer loss. Args: diff --git a/jax3d/projects/generative/nerf/positional_encoding.py b/jax3d/projects/generative/nerf/positional_encoding.py index f9d9fbe..f2b4fc0 100644 --- a/jax3d/projects/generative/nerf/positional_encoding.py +++ b/jax3d/projects/generative/nerf/positional_encoding.py @@ -21,12 +21,12 @@ def sinusoidal( - position: FloatArray, + position: FloatArray, # pyrefly: ignore[not-a-type] minimum_frequency_power: int, maximum_frequency_power: int, include_identity: bool = False, - filter_fn: Optional[Callable[[FloatArray], - FloatArray]] = None) -> FloatArray: + filter_fn: Optional[Callable[[FloatArray], # pyrefly: ignore[not-a-type] + FloatArray]] = None) -> FloatArray: # pyrefly: ignore[not-a-type] """Computes the psotional encoding value from sample positions. Arguments: diff --git a/jax3d/projects/generative/nerf/run_trainer.py b/jax3d/projects/generative/nerf/run_trainer.py index f0095a0..3c7af04 100644 --- a/jax3d/projects/generative/nerf/run_trainer.py +++ b/jax3d/projects/generative/nerf/run_trainer.py @@ -49,7 +49,7 @@ def main(argv): bindings=FLAGS.gin_bindings, skip_unknown=False) - trainer = configs.ExperimentConfig().trainer( + trainer = configs.ExperimentConfig().trainer( # pyrefly: ignore[not-callable] experiment_name=FLAGS.experiment_name, working_dir=FLAGS.base_folder) if FLAGS.mode == "train": diff --git a/jax3d/projects/generative/nerf/trainer.py b/jax3d/projects/generative/nerf/trainer.py index 520ecb6..dec9c41 100644 --- a/jax3d/projects/generative/nerf/trainer.py +++ b/jax3d/projects/generative/nerf/trainer.py @@ -148,7 +148,7 @@ def train_step(self, train_state: TrainState, inputs: Any, scratch: Updated scratch state object. """ del inputs - return train_state.replace(step=train_state.step + 1), None, scratch + return train_state.replace(step=train_state.step + 1), None, scratch # pyrefly: ignore[missing-attribute] def init_state(self) -> TrainState: """Initializes training state.""" @@ -217,7 +217,7 @@ def train(self) -> None: scratch = None collected_summary_data = [] while state.step < self.max_steps: - inputs = next(data_loader) + inputs = next(data_loader) # pyrefly: ignore[bad-argument-type] state, summary_data, scratch = self.train_step(state, inputs, scratch) collected_summary_data.append(summary_data) @@ -236,9 +236,9 @@ def train(self) -> None: self.save_checkpoint(state) checkpoint_state = state - if (state.step - log_state.last_log_step) >= self.log_every: + if (state.step - log_state.last_log_step) >= self.log_every: # pyrefly: ignore[unbound-name] logging.info("Training Step %d / %d", state.step, self.max_steps) - self.write_summaries(summary_writer, collected_summary_data, state.step) + self.write_summaries(summary_writer, collected_summary_data, state.step) # pyrefly: ignore[unbound-name] collected_summary_data = [] log_state = self.write_utilization_summary(summary_writer, state.step, log_state)