Buckets:
| import argparse | |
| try: | |
| from .dataset import prepare_iris | |
| from .model import NeuralNetwork | |
| except (ImportError, ModuleNotFoundError): | |
| from dataset import prepare_iris | |
| from model import NeuralNetwork | |
| def accuracy(pred, true): | |
| return float((pred == true).sum()) / len(true) | |
| def parse_args(): | |
| parser = argparse.ArgumentParser(description="Train a from-scratch Iris neural network.") | |
| parser.add_argument("--hidden-sizes", nargs="+", type=int, default=[32, 16], | |
| help="Sizes for hidden layers.") | |
| parser.add_argument("--epochs", type=int, default=300, help="Number of training epochs.") | |
| parser.add_argument("--batch-size", type=int, default=16, help="Mini-batch size for training.") | |
| parser.add_argument("--lr", type=float, default=0.01, help="Learning rate.") | |
| parser.add_argument("--optimizer", choices=["adam", "momentum", "sgd"], default="adam", | |
| help="Optimizer to use during training.") | |
| parser.add_argument("--dropout", type=float, default=0.0, help="Dropout rate for hidden layers.") | |
| parser.add_argument("--reg", type=float, default=1e-4, help="L2 regularization strength.") | |
| parser.add_argument("--seed", type=int, default=42, help="Random seed.") | |
| return parser.parse_args() | |
| def main(): | |
| args = parse_args() | |
| X_train, X_val, X_test, Y_train, Y_val, Y_test, y_train, y_val, y_test = prepare_iris( | |
| test_size=0.2, val_size=0.1, seed=args.seed | |
| ) | |
| layer_sizes = [X_train.shape[1], *args.hidden_sizes, Y_train.shape[1]] | |
| model = NeuralNetwork( | |
| layer_sizes=layer_sizes, | |
| lr=args.lr, | |
| optimizer=args.optimizer, | |
| reg_lambda=args.reg, | |
| dropout=args.dropout, | |
| seed=args.seed, | |
| ) | |
| model.fit( | |
| X_train, | |
| Y_train, | |
| X_val=X_val, | |
| Y_val=Y_val, | |
| epochs=args.epochs, | |
| batch_size=args.batch_size, | |
| patience=25, | |
| verbose=True, | |
| ) | |
| preds = model.predict(X_test) | |
| acc = accuracy(preds, y_test) | |
| print(f"Test accuracy: {acc * 100:.2f}%") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 2.11 kB
- Xet hash:
- 5e814f74297b133dda242ca49cc33a9615a4d3cf6c2759d507191ab8704a27f2
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.