(name, kwargs)
| 122 | |
| 123 | |
| 124 | def build_algorithm(name, kwargs): |
| 125 | if name == "gradient_clipping": |
| 126 | return algorithms.GradientClipping(**kwargs) |
| 127 | elif name == "alibi": |
| 128 | return algorithms.Alibi(**kwargs) |
| 129 | elif name == "gated_linear_units": |
| 130 | return algorithms.GatedLinearUnits(**kwargs) |
| 131 | elif name == "ema": |
| 132 | return algorithms.EMA( |
| 133 | half_life=kwargs.get("half_life", "1000ba"), |
| 134 | smoothing=kwargs.get("smoothing", None), |
| 135 | ema_start=kwargs.get("ema_start", "0.0dur"), |
| 136 | update_interval=kwargs.get("update_interval", None), |
| 137 | ) |
| 138 | elif name == "rope_schedule": |
| 139 | return FlexBertRopeSchedule( |
| 140 | min_rope_theta=kwargs.get("min_rope_theta", 10_000), |
| 141 | max_rope_theta=kwargs.get("max_rope_theta", 80_000), |
| 142 | warmup_tokens=kwargs.get("warmup_tokens", 25_000_000), |
| 143 | rope_theta_increment=kwargs.get("rope_theta_increment", 10_000), |
| 144 | batch_log_interval=kwargs.get("batch_log_interval", 10), |
| 145 | increment_theta_immediately=kwargs.get("increment_theta_immediately", False), |
| 146 | ) |
| 147 | else: |
| 148 | raise ValueError(f"Not sure how to build algorithm: {name}") |
| 149 | |
| 150 | |
| 151 | def build_callback(name, kwargs): |
no test coverage detected