1 package net.bmahe.genetics4j.gp.program;
2
3 import java.util.Objects;
4 import java.util.random.RandomGenerator;
5
6 import org.apache.commons.lang3.Validate;
7
8 import net.bmahe.genetics4j.core.chromosomes.TreeNode;
9 import net.bmahe.genetics4j.gp.Operation;
10 import net.bmahe.genetics4j.gp.OperationFactory;
11
12 public class StdProgramGenerator implements ProgramGenerator {
13
14 private final ProgramHelper programHelper;
15 private final RandomGenerator randomGenerator;
16
17 @SuppressWarnings({ "unchecked", "rawtypes" })
18 private <T, U> TreeNode<Operation<T>> generate(final Program program, Class<U> acceptedType, int maxDepth,
19 int depth) {
20
21 OperationFactory currentNode = depth < maxDepth - 1 && randomGenerator.nextDouble() < 0.5
22 ? programHelper.pickRandomFunction(program, acceptedType)
23 : programHelper.pickRandomTerminal(program, acceptedType);
24
25 final Operation<T> currentOperation = currentNode.build(program.inputSpec());
26 final TreeNode<Operation<T>> currentTreeNode = new TreeNode<>(currentOperation);
27
28 final Class[] acceptedTypes = currentNode.acceptedTypes();
29
30 for (int i = 0; i < acceptedTypes.length; i++) {
31 final Class childAcceptedType = acceptedTypes[i];
32 final TreeNode<Operation<T>> operation = generate(program, childAcceptedType, maxDepth, depth + 1);
33
34 currentTreeNode.addChild(operation);
35 }
36
37 return currentTreeNode;
38 }
39
40 public StdProgramGenerator(final ProgramHelper _programHelper, final RandomGenerator _randomGenerator) {
41 Objects.requireNonNull(_programHelper);
42 Objects.requireNonNull(_randomGenerator);
43
44 this.programHelper = _programHelper;
45 this.randomGenerator = _randomGenerator;
46 }
47
48 @SuppressWarnings("rawtypes")
49 @Override
50 public TreeNode<Operation<?>> generate(final Program program, final int maxDepth) {
51 Objects.requireNonNull(program);
52 Validate.isTrue(maxDepth > 0);
53
54 final OperationFactory currentNode = randomGenerator.nextDouble() < 0.98
55 ? programHelper.pickRandomFunction(program)
56 : programHelper.pickRandomTerminal(program);
57
58 final Operation currentOperation = currentNode.build(program.inputSpec());
59 final TreeNode<Operation<?>> currentTreeNode = new TreeNode<>(currentOperation);
60
61 Class[] acceptedTypes = currentNode.acceptedTypes();
62
63 for (int i = 0; i < acceptedTypes.length; i++) {
64 final Class acceptedType = acceptedTypes[i];
65 final TreeNode<Operation<?>> operation = generate(program, acceptedType, maxDepth, 1);
66
67 currentTreeNode.addChild(operation);
68 }
69
70 return currentTreeNode;
71 }
72
73 @Override
74 public TreeNode<Operation<?>> generate(final Program program) {
75 return generate(program, program.maxDepth());
76 }
77
78 @Override
79 public <T, U> TreeNode<Operation<T>> generate(final Program program, final int maxDepth, final Class<U> rootType) {
80 Objects.requireNonNull(program);
81 Objects.requireNonNull(rootType);
82 Validate.isTrue(maxDepth > 0);
83
84 return generate(program, rootType, maxDepth, 0);
85 }
86 }