Merge pull request #22 from vmihailenco/master
Add TrimmedMean
- Id
- 47130feefefc1149cea997b1420d0a6c41a77672
- Author
- Caio
- Commit time
- 2018-10-23T13:40:21+02:00
Modified Gopkg.lock
# This file is autogenerated, do not edit; changes may be undone by the next 'dep ensure'.
[[projects]]
+ digest = "1:cf63454c1e81409484ded047413228de0f7a3031f0fcd36d4e1db7620c3c7d1b"
name = "github.com/leesper/go_rng"
packages = ["."]
+ pruneopts = ""
revision = "5344a9259b21627d94279721ab1f27eb029194e7"
[[projects]]
+ digest = "1:ba1b34793902651895b9cd923ea833010bd1497b1bf24f0a65b98db1078475c9"
name = "github.com/yourbasic/fenwick"
packages = ["."]
+ pruneopts = ""
revision = "7cb325001daae879940edf7b19da8557a51ca5f7"
version = "1.2.0"
+
+[[projects]]
+ branch = "master"
+ digest = "1:ad6d9b2cce40c7c44952d49a6a324a2110db43b4279d9e599db74e45de5ae80c"
+ name = "gonum.org/v1/gonum"
+ packages = [
+ "blas",
+ "blas/blas64",
+ "blas/gonum",
+ "floats",
+ "internal/asm/c128",
+ "internal/asm/f32",
+ "internal/asm/f64",
+ "internal/math32",
+ "lapack",
+ "lapack/gonum",
+ "lapack/lapack64",
+ "mat",
+ "stat",
+ ]
+ pruneopts = ""
+ revision = "f0982070f509ee139841ca385c44dc22a77c8da8"
[solve-meta]
analyzer-name = "dep"
analyzer-version = 1
- inputs-digest = "eac5a7bc22a9b8d592de2752be2ab1f8707df0caf1a1e77087c88671cbb8c94d"
+ input-imports = [
+ "github.com/leesper/go_rng",
+ "github.com/yourbasic/fenwick",
+ "gonum.org/v1/gonum/stat",
+ ]
solver-name = "gps-cdcl"
solver-version = 1
Modified tdigest.go
return closest
}
+// TrimmedMean returns the mean of the distribution between the two percentiles
+// p1 and p2.
+func (t *TDigest) TrimmedMean(p1, p2 float64) float64 {
+ if p1 < 0 || p1 > 1 {
+ panic("p1 must be between 0 and 1 (inclusive)")
+ }
+ if p2 < 0 || p2 > 1 {
+ panic("p2 must be between 0 and 1 (inclusive)")
+ }
+ if p1 >= p2 {
+ panic("p1 must be lower than p2")
+ }
+
+ minCount := p1 * float64(t.count)
+ maxCount := p2 * float64(t.count)
+
+ var trimmedSum, trimmedCount, currCount float64
+ for i, mean := range t.summary.means {
+ count := float64(t.summary.counts[i])
+
+ nextCount := currCount + count
+ if nextCount <= minCount {
+ currCount = nextCount
+ continue
+ }
+
+ if currCount < minCount {
+ count = nextCount - minCount
+ }
+ if nextCount > maxCount {
+ count -= nextCount - maxCount
+ }
+
+ trimmedSum += count * mean
+ trimmedCount += count
+
+ if nextCount >= maxCount {
+ break
+ }
+ currCount = nextCount
+ }
+
+ if trimmedCount == 0 {
+ return 0
+ }
+ return trimmedSum / trimmedCount
+}
+
func shuffle(means []float64, counts []uint32, rng RNG) {
for i := len(means) - 1; i > 1; i-- {
j := rng.Intn(i + 1)
Modified tdigest_test.go
"testing"
"github.com/leesper/go_rng"
+ "gonum.org/v1/gonum/stat"
)
func init() {
if cdf := td.CDF(7.144560976650238e+06); cdf > 1 {
t.Fatalf("invalid: %v", cdf)
}
+}
+
+func TestTrimmedMean(t *testing.T) {
+ tests := []struct {
+ p1, p2 float64
+ }{
+ {0, 1},
+ {0.1, 0.9},
+ {0.2, 0.8},
+ {0.25, 0.75},
+ {0, 0.5},
+ {0.5, 1},
+ {0.1, 0.7},
+ {0.3, 0.9},
+ }
+
+ for _, size := range []int{100, 1000, 10000} {
+ for _, test := range tests {
+ td := uncheckedNew(Compression(100))
+
+ data := make([]float64, 0, size)
+ for i := 0; i < size; i++ {
+ f := rand.Float64()
+ data = append(data, f)
+ err := td.Add(f)
+ if err != nil {
+ t.Fatal(err)
+ }
+ }
+
+ got := td.TrimmedMean(test.p1, test.p2)
+ wanted := trimmedMean(data, test.p1, test.p2)
+ if math.Abs(got-wanted) > 0.01 {
+ t.Fatalf("got %f, wanted %f (size=%d p1=%f p2=%f)",
+ got, wanted, size, test.p1, test.p2)
+ }
+
+ for i := 0; i < 10; i++ {
+ err := td.Add(float64(i * 100))
+ if err != nil {
+ t.Fatal(err)
+ }
+ }
+ mean := td.TrimmedMean(0.1, 0.999)
+ if mean < 0 {
+ t.Fatalf("mean < 0")
+ }
+ }
+ }
+}
+
+func TestTrimmedMeanCornerCases(t *testing.T) {
+ td := uncheckedNew(Compression(100))
+
+ mean := td.TrimmedMean(0, 1)
+ if mean != 0 {
+ t.Fatalf("got %f, wanted 0", mean)
+ }
+
+ x := 1.0
+ err := td.Add(x)
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ mean = td.TrimmedMean(0, 1)
+ if mean != 1 {
+ t.Fatalf("got %f, wanted %f", mean, x)
+ }
+
+ err = td.Add(1000)
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ mean = td.TrimmedMean(0, 1)
+ wanted := 500.5
+ if !closeEnough(mean, wanted) {
+ t.Fatalf("got %f, wanted %f", mean, wanted)
+ }
+}
+
+func trimmedMean(ff []float64, p1, p2 float64) float64 {
+ sort.Float64s(ff)
+ x1 := stat.Quantile(p1, stat.Empirical, ff, nil)
+ x2 := stat.Quantile(p2, stat.Empirical, ff, nil)
+
+ var sum float64
+ var count int
+ for _, f := range ff {
+ if f >= x1 && f <= x2 {
+ sum += f
+ count++
+ }
+ }
+ return sum / float64(count)
}
func benchmarkAdd(compression uint32, b *testing.B) {