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