Rixf123's picture
download
raw
2.11 kB
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.