Files
ma/k8s-dvfs/src/main/groovy/recreationaltech/plugin/K8sLinearFitRuntimeEstimator.groovy

100 lines
3.2 KiB
Groovy

package recreationaltech.plugin
import groovy.util.logging.Slf4j
import java.nio.file.Files
import java.nio.file.Path
/**
* Provides runtime estimates for tasks based on task-name and input file size
*/
@Slf4j
class K8sLinearFitRuntimeEstimator extends K8sRuntimeEstimator {
class Function {
private double m
private double n
Function(double m, double n) {
this.m = m
this.n = n
}
double estimate(long x) {
return m * (double)x + n
}
}
HashMap<String, Function> estimators;
/**
* Initialize the estimator with data recorded by K8sRuntimeRecorder
* @param dataFilePath
*/
K8sLinearFitRuntimeEstimator(String dataFilePath) {
def data = parseDataFile(dataFilePath)
fit(data)
}
/**
* Initializes the runtime estimator with statically known data.
* @param data Map from task name to list of recordings, where each recording is a tuple (input-size, runtime-in-ms)
*/
K8sLinearFitRuntimeEstimator(HashMap<String, ArrayList<Tuple2<Long, Long>>> data) {
fit(data)
}
/**
* Returns an estimation of the task runtime in milliseconds
* @param taskName the name of the task
* @param inputSize the total input size in bytes
* @return estimated task runtime in milliseconds OR infinity if the task is unknown.
*/
double estimate(String taskName, long inputSize) {
Function estimator = estimators.get(taskName)
if (estimator == null) {
//log.warn "[K8s] Unable to estimate take ${taskName}: unknown task"
return Double.POSITIVE_INFINITY
}
return estimator.estimate(inputSize)
}
private void fit(HashMap<String, ArrayList<Tuple2<Long, Long>>> data) {
estimators = new HashMap<>()
for (Map.Entry<String, ArrayList<Tuple2<Long, Long>>> t : data) {
Function f = fit(t.value)
estimators.put(t.key, f)
}
}
/**
* Uses linear regression to fit a linear function (y = m * x + n) to the provided observations
* @param observations list of tuples (input size, runtime in ms)
* @return linear function fitted to the input
*/
private Function fit(ArrayList<Tuple2<Long, Long>> observations) throws IllegalArgumentException {
int n = observations.size()
if (n > 1) {
double sumX = 0, sumY = 0, sumXY = 0, sumX2 = 0
for (Tuple2<Long, Long> o : observations) {
double x = (double) o.get(0)
double y = (double) o.get(1)
sumX += x
sumY += y
sumXY += x * y
sumX2 += x * x
}
double m = (n * sumXY - sumX * sumY) / (n * sumX2 - sumX * sumX)
return new Function(
m,
(sumY - m * sumX) / n
)
} else if (n == 1) {
// Special case: We only have 1 measurement. We will just assume that the runtime is constant,
// because in our observed data, it is.
return new Function(0.0, (double) observations[0].get(1))
}
throw new IllegalArgumentException("requires at least 1 observation")
}
}