Skip to content

Commit 593c620

Browse files
riccardopinosioRJKeevil
authored andcommitted
make statistics a proper struct
1 parent 8411fdb commit 593c620

10 files changed

Lines changed: 112 additions & 109 deletions

backends/pipeline.go

Lines changed: 41 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,14 @@
11
package backends
22

33
import (
4+
"encoding/json"
45
"errors"
56
"fmt"
7+
"math"
8+
"time"
69

710
"github.com/knights-analytics/hugot/options"
11+
"github.com/knights-analytics/hugot/util/safeconv"
812
)
913

1014
// BasePipeline can be embedded by a pipeline.
@@ -53,13 +57,49 @@ type PipelineBatchOutput interface {
5357

5458
// Pipeline is the interface that any pipeline must implement.
5559
type Pipeline interface {
56-
GetStats() []string // Get the pipeline running stats
60+
GetStatistics() PipelineStatistics // Get the pipeline running stats
5761
Validate() error // Validate the pipeline for correctness
5862
GetMetadata() PipelineMetadata // Return metadata information for the pipeline
5963
GetModel() *Model // Return the model used by the pipeline
6064
Run([]string) (PipelineBatchOutput, error) // Run the pipeline on an input
6165
}
6266

67+
type PipelineStatistics struct {
68+
TokenizerTotalTime time.Duration
69+
TokenizerExecutionCount uint64
70+
TokenizerAvgQueryTime time.Duration
71+
OnnxTotalTime time.Duration
72+
OnnxExecutionCount uint64
73+
OnnxAvgQueryTime time.Duration
74+
TotalQueries uint64
75+
TotalDocuments uint64
76+
AverageLatency time.Duration
77+
AverageBatchSize float64
78+
FilteredResults uint64
79+
}
80+
81+
func (p *PipelineStatistics) ComputeTokenizerStatistics(timings *timings) {
82+
p.TokenizerTotalTime = safeconv.U64ToDuration(timings.TotalNS)
83+
p.TokenizerExecutionCount = timings.NumCalls
84+
p.TokenizerAvgQueryTime = time.Duration(float64(timings.TotalNS) /
85+
math.Max(1, float64(timings.NumCalls)))
86+
}
87+
88+
func (p *PipelineStatistics) ComputeOnnxStatistics(timings *timings) {
89+
p.OnnxTotalTime = safeconv.U64ToDuration(timings.TotalNS)
90+
p.OnnxExecutionCount = timings.NumCalls
91+
p.OnnxAvgQueryTime = time.Duration(float64(timings.TotalNS) /
92+
math.Max(1, float64(timings.NumCalls)))
93+
}
94+
95+
func (p *PipelineStatistics) Print() {
96+
jsonData, err := json.MarshalIndent(p, "", " ")
97+
if err != nil {
98+
fmt.Println(err)
99+
}
100+
fmt.Println(string(jsonData))
101+
}
102+
63103
// PipelineOption is an option for a pipeline type.
64104
type PipelineOption[T Pipeline] func(eo T) error
65105

hugot.go

Lines changed: 21 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -58,10 +58,10 @@ func newSession(backend string, opts ...options.WithOption) (*Session, error) {
5858

5959
type pipelineMap[T backends.Pipeline] map[string]T
6060

61-
func (m pipelineMap[T]) GetStats() []string {
62-
var stats []string
61+
func (m pipelineMap[T]) GetStatistics() []backends.PipelineStatistics {
62+
var stats []backends.PipelineStatistics
6363
for _, p := range m {
64-
stats = append(stats, p.GetStats()...)
64+
stats = append(stats, p.GetStatistics())
6565
}
6666
return stats
6767
}
@@ -379,24 +379,33 @@ func (e *pipelineNotFoundError) Error() string {
379379
return fmt.Sprintf("Pipeline with name %s not found", e.pipelineName)
380380
}
381381

382-
// GetStats returns runtime statistics for all initialized pipelines for profiling purposes. We currently record for each pipeline:
382+
// GetStatistics returns runtime statistics for all initialized pipelines for profiling purposes. We currently record for each pipeline:
383383
// the total runtime of the tokenization step
384384
// the number of batch calls to the tokenization step
385385
// the average time per tokenization batch call
386386
// the total runtime of the inference (i.e. onnxruntime) step
387387
// the number of batch calls to the onnxruntime inference
388388
// the average time per onnxruntime inference batch call.
389-
func (s *Session) GetStats() []string {
389+
func (s *Session) GetStatistics() []backends.PipelineStatistics {
390390
return slices.Concat(
391-
s.tokenClassificationPipelines.GetStats(),
392-
s.textClassificationPipelines.GetStats(),
393-
s.featureExtractionPipelines.GetStats(),
394-
s.zeroShotClassificationPipelines.GetStats(),
395-
s.crossEncoderPipelines.GetStats(),
396-
s.textGenerationPipelines.GetStats(),
391+
s.tokenClassificationPipelines.GetStatistics(),
392+
s.textClassificationPipelines.GetStatistics(),
393+
s.featureExtractionPipelines.GetStatistics(),
394+
s.imageClassificationPipelines.GetStatistics(),
395+
s.zeroShotClassificationPipelines.GetStatistics(),
396+
s.crossEncoderPipelines.GetStatistics(),
397+
s.textGenerationPipelines.GetStatistics(),
397398
)
398399
}
399400

401+
// Print prints runtime statistics for all initialized pipelines to stdout.
402+
func (s *Session) PrintStatistics() {
403+
stats := s.GetStatistics()
404+
for _, v := range stats {
405+
v.Print()
406+
}
407+
}
408+
400409
// Destroy deletes the hugot session and onnxruntime environment and all initialized pipelines, freeing memory.
401410
// A hugot session should be destroyed when not neeeded any more, preferably with a defer() call.
402411
func (s *Session) Destroy() error {
@@ -408,6 +417,7 @@ func (s *Session) Destroy() error {
408417
s.featureExtractionPipelines = nil
409418
s.tokenClassificationPipelines = nil
410419
s.textClassificationPipelines = nil
420+
s.imageClassificationPipelines = nil
411421
s.zeroShotClassificationPipelines = nil
412422
s.textGenerationPipelines = nil
413423
s.crossEncoderPipelines = nil

hugot_test.go

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -236,8 +236,8 @@ func textClassificationPipeline(t *testing.T, session *Session) {
236236
}
237237
})
238238

239-
// check get stats
240-
session.GetStats()
239+
// check print stats
240+
session.PrintStatistics()
241241
}
242242

243243
func textClassificationPipelineMulti(t *testing.T, session *Session) {
@@ -396,7 +396,10 @@ func textClassificationPipelineMulti(t *testing.T, session *Session) {
396396
})
397397

398398
// check get stats
399-
session.GetStats()
399+
stats := session.GetStatistics()
400+
for _, v := range stats {
401+
v.Print()
402+
}
400403
}
401404

402405
func textClassificationPipelineValidation(t *testing.T, session *Session) {

pipelines/crossEncoder.go

Lines changed: 10 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ package pipelines
33
import (
44
"errors"
55
"fmt"
6-
"math"
76
"sort"
87
"sync/atomic"
98
"time"
@@ -104,27 +103,20 @@ func (p *CrossEncoderPipeline) GetMetadata() backends.PipelineMetadata {
104103
}
105104
}
106105

107-
func (p *CrossEncoderPipeline) GetStats() []string {
106+
func (p *CrossEncoderPipeline) GetStatistics() backends.PipelineStatistics {
108107
avgLatency := p.stats.AverageLatency
109108
if p.stats.TotalQueries > 0 {
110109
avgLatency = time.Duration(float64(p.stats.AverageLatency) / float64(p.stats.TotalQueries))
111110
}
112-
return []string{
113-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
114-
fmt.Sprintf("Total queries processed: %d", p.stats.TotalQueries),
115-
fmt.Sprintf("Total documents scored: %d", p.stats.TotalDocuments),
116-
fmt.Sprintf("Average latency per query: %s", avgLatency),
117-
fmt.Sprintf("Average batch size: %.2f", p.stats.AverageBatchSize),
118-
fmt.Sprintf("Filtered results: %d", p.stats.FilteredResults),
119-
fmt.Sprintf("Tokenizer: Total time=%s, Execution count=%d, Average query time=%s",
120-
safeconv.U64ToDuration(p.Model.Tokenizer.TokenizerTimings.TotalNS),
121-
p.Model.Tokenizer.TokenizerTimings.NumCalls,
122-
time.Duration(float64(p.Model.Tokenizer.TokenizerTimings.TotalNS)/math.Max(1, float64(p.Model.Tokenizer.TokenizerTimings.NumCalls)))),
123-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
124-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
125-
p.PipelineTimings.NumCalls,
126-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
127-
}
111+
statistics := backends.PipelineStatistics{}
112+
statistics.ComputeTokenizerStatistics(p.Model.Tokenizer.TokenizerTimings)
113+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
114+
statistics.TotalQueries = p.stats.TotalQueries
115+
statistics.TotalDocuments = p.stats.TotalDocuments
116+
statistics.AverageLatency = avgLatency
117+
statistics.AverageBatchSize = p.stats.AverageBatchSize
118+
statistics.FilteredResults = p.stats.FilteredResults
119+
return statistics
128120
}
129121

130122
func (p *CrossEncoderPipeline) Validate() error {

pipelines/featureExtraction.go

Lines changed: 6 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ package pipelines
33
import (
44
"errors"
55
"fmt"
6-
"math"
76
"strings"
87
"sync/atomic"
98
"time"
@@ -112,19 +111,12 @@ func (p *FeatureExtractionPipeline) GetMetadata() backends.PipelineMetadata {
112111
}
113112
}
114113

115-
// GetStats returns the runtime statistics for the pipeline.
116-
func (p *FeatureExtractionPipeline) GetStats() []string {
117-
return []string{
118-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
119-
fmt.Sprintf("Tokenizer: Total time=%s, Execution count=%d, Average query time=%s",
120-
safeconv.U64ToDuration(p.Model.Tokenizer.TokenizerTimings.TotalNS),
121-
p.Model.Tokenizer.TokenizerTimings.NumCalls,
122-
time.Duration(float64(p.Model.Tokenizer.TokenizerTimings.TotalNS)/math.Max(1, float64(p.Model.Tokenizer.TokenizerTimings.NumCalls)))),
123-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
124-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
125-
p.PipelineTimings.NumCalls,
126-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
127-
}
114+
// GetStatistics returns the runtime statistics for the pipeline.
115+
func (p *FeatureExtractionPipeline) GetStatistics() backends.PipelineStatistics {
116+
statistics := backends.PipelineStatistics{}
117+
statistics.ComputeTokenizerStatistics(p.Model.Tokenizer.TokenizerTimings)
118+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
119+
return statistics
128120
}
129121

130122
// Validate checks that the pipeline is valid.

pipelines/imageClassification.go

Lines changed: 4 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@ import (
66
"image"
77
_ "image/jpeg" // adds jpeg support to image
88
_ "image/png" // adds png support to image
9-
"math"
109
"sort"
1110
"strings"
1211
"sync/atomic"
@@ -147,14 +146,10 @@ func (p *ImageClassificationPipeline) GetMetadata() backends.PipelineMetadata {
147146
}
148147
}
149148

150-
func (p *ImageClassificationPipeline) GetStats() []string {
151-
return []string{
152-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
153-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
154-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
155-
p.PipelineTimings.NumCalls,
156-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
157-
}
149+
func (p *ImageClassificationPipeline) GetStatistics() backends.PipelineStatistics {
150+
statistics := backends.PipelineStatistics{}
151+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
152+
return statistics
158153
}
159154

160155
func (p *ImageClassificationPipeline) Validate() error {

pipelines/textClassification.go

Lines changed: 6 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ package pipelines
33
import (
44
"errors"
55
"fmt"
6-
"math"
76
"sync/atomic"
87
"time"
98

@@ -129,19 +128,12 @@ func (p *TextClassificationPipeline) GetMetadata() backends.PipelineMetadata {
129128
}
130129
}
131130

132-
// GetStats returns the runtime statistics for the pipeline.
133-
func (p *TextClassificationPipeline) GetStats() []string {
134-
return []string{
135-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
136-
fmt.Sprintf("Tokenizer: Total time=%s, Execution count=%d, Average query time=%s",
137-
safeconv.U64ToDuration(p.Model.Tokenizer.TokenizerTimings.TotalNS),
138-
p.Model.Tokenizer.TokenizerTimings.NumCalls,
139-
time.Duration(float64(p.Model.Tokenizer.TokenizerTimings.TotalNS)/math.Max(1, float64(p.Model.Tokenizer.TokenizerTimings.NumCalls)))),
140-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
141-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
142-
p.PipelineTimings.NumCalls,
143-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
144-
}
131+
// GetStatistics returns the runtime statistics for the pipeline.
132+
func (p *TextClassificationPipeline) GetStatistics() backends.PipelineStatistics {
133+
statistics := backends.PipelineStatistics{}
134+
statistics.ComputeTokenizerStatistics(p.Model.Tokenizer.TokenizerTimings)
135+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
136+
return statistics
145137
}
146138

147139
// Validate checks that the pipeline is valid.

pipelines/textGeneration.go

Lines changed: 6 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,6 @@ import (
44
"bytes"
55
"errors"
66
"fmt"
7-
"math"
87
"sync/atomic"
98
"text/template"
109
"time"
@@ -167,18 +166,12 @@ func (p *TextGenerationPipeline) GetModel() *backends.Model {
167166
return p.Model
168167
}
169168

170-
func (p *TextGenerationPipeline) GetStats() []string {
171-
return []string{
172-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
173-
fmt.Sprintf("Tokenizer: Total time=%s, Execution count=%d, Average query time=%s",
174-
safeconv.U64ToDuration(p.Model.Tokenizer.TokenizerTimings.TotalNS),
175-
p.Model.Tokenizer.TokenizerTimings.NumCalls,
176-
time.Duration(float64(p.Model.Tokenizer.TokenizerTimings.TotalNS)/math.Max(1, float64(p.Model.Tokenizer.TokenizerTimings.NumCalls)))),
177-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
178-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
179-
p.PipelineTimings.NumCalls,
180-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
181-
}
169+
// GetStatistics returns the runtime statistics for the pipeline.
170+
func (p *TextGenerationPipeline) GetStatistics() backends.PipelineStatistics {
171+
statistics := backends.PipelineStatistics{}
172+
statistics.ComputeTokenizerStatistics(p.Model.Tokenizer.TokenizerTimings)
173+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
174+
return statistics
182175
}
183176

184177
func (p *TextGenerationPipeline) Validate() error {

pipelines/tokenClassification.go

Lines changed: 6 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ package pipelines
33
import (
44
"errors"
55
"fmt"
6-
"math"
76
"slices"
87
"strings"
98
"sync/atomic"
@@ -149,19 +148,12 @@ func (p *TokenClassificationPipeline) GetMetadata() backends.PipelineMetadata {
149148
}
150149
}
151150

152-
// GetStats returns the runtime statistics for the pipeline.
153-
func (p *TokenClassificationPipeline) GetStats() []string {
154-
return []string{
155-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
156-
fmt.Sprintf("Tokenizer: Total time=%s, Execution count=%d, Average query time=%s",
157-
safeconv.U64ToDuration(p.Model.Tokenizer.TokenizerTimings.TotalNS),
158-
p.Model.Tokenizer.TokenizerTimings.NumCalls,
159-
time.Duration(float64(p.Model.Tokenizer.TokenizerTimings.TotalNS)/math.Max(1, float64(p.Model.Tokenizer.TokenizerTimings.NumCalls)))),
160-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
161-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
162-
p.PipelineTimings.NumCalls,
163-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
164-
}
151+
// GetStatistics returns the runtime statistics for the pipeline.
152+
func (p *TokenClassificationPipeline) GetStatistics() backends.PipelineStatistics {
153+
statistics := backends.PipelineStatistics{}
154+
statistics.ComputeTokenizerStatistics(p.Model.Tokenizer.TokenizerTimings)
155+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
156+
return statistics
165157
}
166158

167159
// Validate checks that the pipeline is valid.

pipelines/zeroShotClassification.go

Lines changed: 6 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -311,18 +311,12 @@ func (p *ZeroShotClassificationPipeline) GetMetadata() backends.PipelineMetadata
311311
}
312312
}
313313

314-
func (p *ZeroShotClassificationPipeline) GetStats() []string {
315-
return []string{
316-
fmt.Sprintf("Statistics for pipeline: %s", p.PipelineName),
317-
fmt.Sprintf("Tokenizer: Total time=%s, Execution count=%d, Average query time=%s",
318-
safeconv.U64ToDuration(p.Model.Tokenizer.TokenizerTimings.TotalNS),
319-
p.Model.Tokenizer.TokenizerTimings.NumCalls,
320-
time.Duration(float64(p.Model.Tokenizer.TokenizerTimings.TotalNS)/math.Max(1, float64(p.Model.Tokenizer.TokenizerTimings.NumCalls)))),
321-
fmt.Sprintf("ONNX: Total time=%s, Execution count=%d, Average query time=%s",
322-
safeconv.U64ToDuration(p.PipelineTimings.TotalNS),
323-
p.PipelineTimings.NumCalls,
324-
time.Duration(float64(p.PipelineTimings.TotalNS)/math.Max(1, float64(p.PipelineTimings.NumCalls)))),
325-
}
314+
// GetStatistics returns the runtime statistics for the pipeline.
315+
func (p *ZeroShotClassificationPipeline) GetStatistics() backends.PipelineStatistics {
316+
statistics := backends.PipelineStatistics{}
317+
statistics.ComputeTokenizerStatistics(p.Model.Tokenizer.TokenizerTimings)
318+
statistics.ComputeOnnxStatistics(p.PipelineTimings)
319+
return statistics
326320
}
327321

328322
func (p *ZeroShotClassificationPipeline) Run(inputs []string) (backends.PipelineBatchOutput, error) {

0 commit comments

Comments
 (0)