(params, key, scale, group_size=CQ_GROUP_SIZE)
| 267 | |
| 268 | |
| 269 | def noise_params(params, key, scale, group_size=CQ_GROUP_SIZE): |
| 270 | return _map_quant_leaves( |
| 271 | params, |
| 272 | lambda w, i: add_cq_noise(w, jax.random.fold_in(key, i), scale, group_size), |
| 273 | ) |
| 274 | |
| 275 | |
| 276 | def noise_params_only(params, key, scale, name, group_size=CQ_GROUP_SIZE): |
nothing calls this directly
no test coverage detected